1 条题解
-
0
P9480 [NOI2023] 深搜
超级酷的题目!
思路参考了 Rainbow_qwq 的题解,对他的题解进行一些补充。题解最后给出的代码有: 的性质 B,,未卡常正解和卡常后的正解。
首先, 可以作为 -DFS 树的充要条件是所有非树边均为返祖边,因为 DFS 树的非树边为返祖边,且如果满足条件则对 进行 DFS 就是符合条件的 DFS 顺序。
称非树边 覆盖 点 当且仅当 在以 为根的 上是横叉边(非返祖边)。易知非树边 不覆盖 点 当且仅当 在以 为根时 的子树内(称为 侧),或以 为根时 的子树内(称为 侧)。
对于性质 B,有两种思路:
- 记录从子树内向上延伸的返祖边的最浅深度,并维护当前子树内是否有关键点未被覆盖。
- 设 表示不覆盖 中任意关键点的非树边数量,设关键点集合为 ,容斥得 $\sum_{|S| \geq 1 \ \land \ S\subseteq U} (-1) ^ {|S| + 1} 2 ^ {c(S)}$。
顺着前一种思路不容易解决原问题,因此考虑第二种思路。
非树边的形态
- 称一个点 属于虚树,当且仅当它是虚树的结点。
- 称一个点 落在虚树上,当且仅当它被虚树的某条边覆盖。
定点 为根。
建出 的虚树 。注意要和我们所理解的虚树做区分:当且仅当 或有至少三棵子树(包括向上的子树)含属于 的点时, 才属于虚树。它们的唯一区别在于,当 所有点的 LCA 只有两棵子树含属于 的点时, 不属于 ,但 属于我们通常所理解的虚树。
考虑非树边 不覆盖 的充要条件: 全部在 侧, 全部在 侧, 侧均有属于 的点。它将合法非树边的形态分成两类:
- 落在虚树上的点向虚树外延伸的某条边的对应子树内的返祖边。注意 “返祖边” 是在以该点为根时判定的。
- 对应树上路径被虚树的某条边完全包含。

如上图,绿色的点表示 ,黄色的点不属于 但属于虚树,橙色的点不属于虚树但落在虚树上。红色虚边为合法非树边,蓝色虚边为不合法非树边。
性质 B
对于性质 B,注意到不存在跨过 的合法横叉边,因此可以将虚树改为我们通常所理解的虚树:若 有两棵子树含属于 的点,则 也属于虚树。即认为 属于虚树对答案无影响。
枚举 的 LCA ,则不存在一端在 子树内(不含 ),另一端在 子树外的(不含 )的合法非树边。因此子树内外独立,贡献直接相乘。
- 子树内的贡献可利用虚树结构 DP,因此设 表示当 属于虚树 时,仅考虑 子树内所有关键点和非树边的贡献。
- 子树外的贡献为 ,其中 表示 子树外以 为根时的返祖边数量。设 表示 子树内以 为根时的返祖边数量,换根 DP 求出 和 。 的求法稍有些复杂,留给读者自行思考,可以参考代码
dfs3部分。
的转移是重头戏。转移时,需要不重不漏地统计每一处贡献。为此,请读者牢记所有可能产生贡献的非树边形态:落在虚树上的点向虚树外延伸的某条边的对应子树内的返祖边;返祖边两端对应树上路径被虚树的某条边完全包含。
设 的所有儿子为 。 的子树内要么没有虚树上的点,要么有虚树上的点。对于后者,根据虚树的性质,有且仅有最浅的点与 在虚树上相连。
设 在虚树上的儿子为 ,有基本贡献系数 。接下来统计未考虑到的非树边:
- 第一类非树边:对于每个 ,包含于路径 的非树边。
- 第二类非树边:对于每个 和所有 (不含 ),从 向虚树外延伸的某条边的对应子树内的返祖边。
- 第三类非树边:对于 ,向虚树外延伸且在子树内的某条边的对应子树内的返祖边。
将第三类非树边摊到每个未选点的儿子,得:
- 对于选点的儿子 ,贡献系数为 $X(s_i) = \sum_{p\in \mathrm{subtree}(s_i)} g_{p\to d}$,其中 表示 乘以 ,其中 表示包含于路径 的非树边数量,加上对于所有 (不含 ),从 向路径 外延伸的某条边的对应子树的返祖边数量之和。
- 对于未选点的儿子 ,贡献系数为 ,其中 表示 子树内的返祖边数量 ,加上较浅的一端为 ,较深的一端落在 子树内的返祖边数量。

如上图, 的左子树选择点 ,红色虚边为产生贡献的非树边。 的右子树没有选择,红色虚边为产生贡献的非树边。
假设已经求出 和 ,考虑 如何计算。设恰好在 棵子树内选点的答案为 ,即 。
- 如果 ,则至少要有两棵子树被选择,产生贡献 。
- 如果 ,则要求 是关键点,对被选择子树数量没有限制,产生贡献 。
发现 的 等价,因此背包时维护 和 即可。
设 ,则 $Z = \sum_{|S|\geq 1 \ \land \ S\subseteq U} (-1) ^ {|S|} 2 ^ {c(S)}$,答案为 。
考虑维护 求 。根据 的定义,可得如下步骤:
- 首先,对于 ,令 ,其中 是 对应的 的儿子。
- 对于所有 ,令 。由于不存在等于树边的非树边,将这一步放在下面两步之后也正确。
- 加入第一类非树边的贡献:枚举所有以 为一端,另一端在 子树内的非树边 ,对于 ,将 乘以 。
- 加入第二类非树边的贡献:枚举所有以 为一端,另一端在 子树内的非树边 ,设 对应 的儿子为 ,则对于 的所有不为 的儿子 ,对于 ,将 乘以 。简单地说,返祖边 (其中 是较浅端)会对 的不为 后继的所有儿子的子树内所有结点产生贡献。
我使用的维护方法是:继承 ;加入第一类非树边的贡献;计算 求出 ;加入以 为较浅端的第二类非树边的贡献;令 。其实就是将统计第二类非树边的贡献下放到每个儿子处进行,这样好写一点,但注意这类贡献需要在计算 求出 之后统计。
将 这一维用 DFS 序拍平到序列上,线段树合并维护 ,支持区间乘法,区间求和。注意到每个子树的时间戳区间不交,所以不用线段树合并,只需用一棵线段树维护。
考虑求 。在计算 ,加入第一类非树边的贡献时,对于每个 以及对应 的儿子 ,将 加上 ,过程结束后得到新的 。则 。
时间复杂度 。
正解
相较于性质 B 多出了以 为根时的横叉边。
考虑在何种情况下性质 B 的做法会导致错误:性质 B 保证我们可以将 加入虚树,而性质 B 的做法在 “无论非树边是否是返祖边,只要 在虚树上” 时,均可得到正确的答案,因为我们在分析合法非树边形态时,并没有要求这些非树边是以 为根时的返祖边。这说明性质 B 只是 让我们将 加入虚树。
将 加入虚树会导致一些非树边从合法变成不合法。
- 对于形如 “落在虚树上的点向虚树外延伸的某条边的对应子树内的返祖边” 的非树边,由于将 加入虚树不改变落在虚树的点集,因此不会有影响。
- 对于形如 “ 对应树上路径被虚树的某条边完全包含” 的非树边,将 加入虚树后产生的影响为:若虚树原本不包含 ,且存在经过 的虚树边 ,那么一条两端均属于 树上路径,且两端分别在 的两侧的非树边从合法变成了不合法。
因此,只需重新计算这样的 的贡献:,且 恰有 两棵子树有属于 的点。
枚举 和将 加入虚树后它在虚树上的两个儿子 ,设它们分别对应 在原树上的儿子 。设 $z = g_{p\to d} \times g_{q\to d}\times \prod_{s_i \neq u, v} Y(s_i)$,这是原来计算的贡献。枚举 LCA 为 的非树边 ,若 分别在 和 上(路径不含 , 顺序无关),则将 乘以 。最终得到 为真正的贡献。暴力做的复杂度是 。
先将 子树时间戳区间的所有 值乘以 ,最后将贡献乘以 (我们发现,实际上 ),这样贡献只与 和 有关,而与其它儿子无关。
枚举 ,则一条 LCA 为 的非树边的贡献是矩形乘以 。若干次矩形乘以 之后全平面求和,注意容斥掉 的贡献。我使用的维护方法是:支持查找子树内某点 对应的 的儿子在时间戳上的后继 。扫描线,将所有儿子的时间戳加入事件,则遇到事件 时(当前扫到时间戳 ,操作为给 乘以 ),设上一次考虑的时间戳为 (初始值为 的儿子时间戳最小值 )。因为所有儿子的时间戳也被加入了事件,所以 一定是同一个儿子子树内的时间戳。求出 的 值之和,求出 的 值之和( 表示 的子树时间戳最大值 ),相乘后再乘以 加入答案。
最后不要忘记将答案减去 ,也就是减去原先计算的错误答案。为此,需要区分 和 ,背包时维护 。
时间复杂度 ,空间复杂度为 ,加上求 LCA 的空间复杂度。代码。
的性质 B
#include <bits/stdc++.h> using namespace std; using ll = long long; using pii = pair<int, int>; using pll = pair<ll, ll>; using pdi = pair<double, int>; using pdd = pair<double, double>; using ull = unsigned long long; #define ppc(x) __builtin_popcount(x) #define clz(x) __builtin_clz(x) bool Mbe; // mt19937 rnd(chrono::steady_clock::now().time_since_epoch().count()); mt19937_64 rnd(1064); int rd(int l, int r) { return rnd() % (r - l + 1) + l; } constexpr int mod = 1e9 + 7; void addt(int &x, int y) { x += y, x >= mod && (x -= mod); } int add(int x, int y) { return x += y, x >= mod && (x -= mod), x; } int ksm(int a, int b) { int s = 1; while(b) { if(b & 1) s = 1ll * s * a % mod; a = 1ll * a * a % mod, b >>= 1; } return s; } constexpr int Z = 1e6 + 5; int fc[Z], ifc[Z]; int bin(int n, int m) { if(n < m) return 0; return 1ll * fc[n] * ifc[m] % mod * ifc[n - m] % mod; } void init_fac(int Z) { for(int i = fc[0] = 1; i < Z; i++) fc[i] = 1ll * fc[i - 1] * i % mod; ifc[Z - 1] = ksm(fc[Z - 1], mod - 2); for(int i = Z - 2; ~i; i--) ifc[i] = 1ll * ifc[i + 1] * (i + 1) % mod; } void cmin(ll &x, ll y) { x = x < y ? x : y; } // ---------- templates above ---------- constexpr int K = 20; constexpr int N = 5000 + 5; constexpr int debug = 0; int n, m, k, vis[N]; vector<int> e[N], son[N], buc[N]; int dn, dfn[N], sz[N], mi[K][N]; bool cmp(int x, int y) {return dfn[x] < dfn[y];} int get(int x, int y) {return dfn[x] < dfn[y] ? x : y;} bool anc(int x, int y) {return dfn[x] <= dfn[y] && dfn[y] < dfn[x] + sz[x];} int suc(int x, int y) { if(!anc(x, y)) return mi[0][dfn[x]]; int l = 0, r = son[x].size() - 1; while(l < r) { int m = l + r + 2 >> 1; if(dfn[y] >= dfn[son[x][m]]) l = m; else r = m - 1; } return son[x][l]; } int lca(int x, int y) { if(x == y) return x; if((x = dfn[x]) > (y = dfn[y])) swap(x, y); int d = __lg(y - x++); return get(mi[d][x], mi[d][y - (1 << d) + 1]); } void dfs1(int id, int ff) { sz[id] = 1, mi[0][dfn[id] = ++dn] = ff; for(int it : e[id]) { if(it == ff) continue; dfs1(it, id), sz[id] += sz[it], son[id].push_back(it); } sort(son[id].begin(), son[id].end(), cmp); } int in[N], out[N]; void dfs2(int id) { for(int it : son[id]) dfs2(it), in[id] += in[it]; for(int it : buc[id]) in[id] += anc(id, it); } void dfs3(int id) { int tot = 0; for(int it : buc[id]) { if(!anc(id, it)) out[id]++; else out[suc(id, it)]--, tot++; } int s = out[id]; for(int it : son[id]) s += in[it]; for(int it : son[id]) out[it] += s - in[it] + tot, dfs3(it); if(debug) { cout << id << " in[id] = " << in[id] << " out[id] = " << out[id] << "\n"; } } int ans, f[N]; void dfs4(int id) { for(int it : son[id]) dfs4(it); if(debug) { cout << "----------------------------------- dfs4 id = " << id << "\n"; } for(int it : buc[id]) { if(!anc(id, it)) continue; int l = dfn[it], r = l + sz[it]; for(int p = l; p < r; p++) addt(f[p], f[p]); in[suc(id, it)]++; } vector<int> F(3); F[0] = 1; for(int it : son[id]) { int pick = 0, np = ksm(2, in[it]); int l = dfn[it], r = l + sz[it]; for(int p = l; p < r; p++) addt(pick, f[p]); vector<int> G(3); if(debug) { cout << "it = " << it << " pick = " << pick << " non pick = " << np << "\n"; } for(int i = 0; i < 3; i++) { addt(G[i], 1ll * F[i] * np % mod); addt(G[min(2, i + 1)], 1ll * F[i] * pick % mod); } F = G; } int pos = dfn[id]; for(int it : son[id]) { int l = dfn[it], r = l + sz[it], c = ksm(2, in[it]); for(int p = pos + 1; p < l; p++) f[p] = 1ll * f[p] * c % mod; for(int p = r; p < pos + sz[id]; p++) f[p] = 1ll * f[p] * c % mod; } if(vis[id]) f[pos] = mod - add(F[0], F[1]); else f[pos] = F[2]; addt(ans, 1ll * f[pos] * ksm(2, out[id]) % mod); if(debug) { cout << id << " pos = " << pos << " vis " << vis[id] << "\n"; for(int i = 1; i <= n; i++) cout << f[i] << " "; cout << "\n"; cout << "F: " << F[0] << " " << F[1] << " " << F[2] << "\n"; cout << "curans = " << ans << "\n"; } } void mian() { cin >> n >> m >> k; for(int i = 1, a, b; i < n; i++) { cin >> a >> b; e[a].push_back(b); e[b].push_back(a); } dfs1(1, 0); for(int i = 1; i <= __lg(n); i++) { for(int j = 1; j + (1 << i) - 1 <= n; j++) { mi[i][j] = get(mi[i - 1][j], mi[i - 1][j + (1 << i - 1)]); } } for(int i = 1, a, b; i <= m; i++) { cin >> a >> b; buc[a].push_back(b); buc[b].push_back(a); } for(int i = 1, a; i <= k; i++) cin >> a, vis[a] = 1; dfs2(1), dfs3(1), dfs4(1); if(debug) { cout << "ans = " << ans << "\n"; } cout << add(mod, -ans) << "\n"; } bool Med; int main() { // freopen("dfs4.in", "r", stdin); // freopen("1.out", "w", stdout); fprintf(stderr, "%.3lf MB\n", (&Mbe - &Med) / 1048576.0); ios::sync_with_stdio(0), cin.tie(0), cout.tie(0); int c, T = 1; cin >> c; while(T--) mian(); cerr << 1e3 * clock() / CLOCKS_PER_SEC << " ms\n"; return 0; }的正解
#include <bits/stdc++.h> using namespace std; using ll = long long; using pii = pair<int, int>; using pll = pair<ll, ll>; using pdi = pair<double, int>; using pdd = pair<double, double>; using ull = unsigned long long; #define ppc(x) __builtin_popcount(x) #define clz(x) __builtin_clz(x) bool Mbe; // mt19937 rnd(chrono::steady_clock::now().time_since_epoch().count()); mt19937_64 rnd(1064); int rd(int l, int r) { return rnd() % (r - l + 1) + l; } constexpr int mod = 1e9 + 7; void addt(int &x, int y) { x += y, x >= mod && (x -= mod); } int add(int x, int y) { return x += y, x >= mod && (x -= mod), x; } int ksm(int a, int b) { int s = 1; while(b) { if(b & 1) s = 1ll * s * a % mod; a = 1ll * a * a % mod, b >>= 1; } return s; } constexpr int Z = 1e6 + 5; int fc[Z], ifc[Z]; int bin(int n, int m) { if(n < m) return 0; return 1ll * fc[n] * ifc[m] % mod * ifc[n - m] % mod; } void init_fac(int Z) { for(int i = fc[0] = 1; i < Z; i++) fc[i] = 1ll * fc[i - 1] * i % mod; ifc[Z - 1] = ksm(fc[Z - 1], mod - 2); for(int i = Z - 2; ~i; i--) ifc[i] = 1ll * ifc[i + 1] * (i + 1) % mod; } void cmin(ll &x, ll y) { x = x < y ? x : y; } // ---------- templates above ---------- constexpr int K = 20; constexpr int N = 5000 + 5; constexpr int debug = 0; int n, m, k, vis[N]; vector<pii> sub[N]; vector<int> e[N], son[N], buc[N]; int dn, dfn[N], rev[N], sz[N], mi[K][N]; bool cmp(int x, int y) { return dfn[x] < dfn[y]; } int get(int x, int y) { return dfn[x] < dfn[y] ? x : y; } bool anc(int x, int y) { return dfn[x] <= dfn[y] && dfn[y] < dfn[x] + sz[x]; } int suc(int x, int y) { if(!anc(x, y)) return mi[0][dfn[x]]; int l = 0, r = son[x].size() - 1; while(l < r) { int m = l + r + 2 >> 1; if(dfn[y] >= dfn[son[x][m]]) l = m; else r = m - 1; } return son[x][l]; } int lca(int x, int y) { if(x == y) return x; if((x = dfn[x]) > (y = dfn[y])) swap(x, y); int d = __lg(y - x++); return get(mi[d][x], mi[d][y - (1 << d) + 1]); } void dfs1(int id, int ff) { sz[id] = 1; mi[0][dfn[id] = ++dn] = ff; rev[dn] = id; for(int it : e[id]) { if(it == ff) continue; dfs1(it, id); sz[id] += sz[it]; son[id].push_back(it); } sort(son[id].begin(), son[id].end(), cmp); } int in[N], out[N]; void dfs2(int id) { for(int it : son[id]) { dfs2(it), in[id] += in[it]; } for(int it : buc[id]) { in[id] += anc(id, it); } } void dfs3(int id) { int tot = 0; for(int it : buc[id]) { if(!anc(id, it)) out[id]++; else out[suc(id, it)]--, tot++; } int s = out[id]; for(int it : son[id]) s += in[it]; for(int it : son[id]) { out[it] += s - in[it] + tot, dfs3(it); } if(debug) { cout << id << " in[id] = " << in[id] << " out[id] = " << out[id] << "\n"; } } int ans, f[N]; void dfs4(int id) { for(int it : son[id]) dfs4(it); if(debug) { cout << "----------------------------------- dfs4 id = " << id << "\n"; } for(int it : buc[id]) { if(!anc(id, it)) continue; int l = dfn[it], r = l + sz[it]; for(int p = l; p < r; p++) addt(f[p], f[p]); in[suc(id, it)]++; } for(int i = 0; i < son[id].size(); i++) { for(int j = i + 1; j < son[id].size(); j++) { int u = son[id][i], v = son[id][j]; int l1 = dfn[u], r1 = l1 + sz[u]; int l2 = dfn[v], r2 = l2 + sz[v]; for(int p = l1; p < r1; p++) { for(int q = l2; q < r2; q++) { int val = 1ll * f[p] * f[q] % mod; for(pii I : sub[id]) { int x = I.first, y = I.second; if(!anc(x, rev[p]) && !anc(y, rev[p])) continue; if(!anc(x, rev[q]) && !anc(y, rev[q])) continue; addt(val, val); } for(int it : son[id]) { if(it != u && it != v) { val = 1ll * val * ksm(2, in[it]) % mod; } } val = 1ll * val * ksm(2, out[id]) % mod; if(debug) { cout << "(p, q) = " << p << ", " << q << "\n"; cout << "val = " << val << "\n"; } addt(ans, val); } } } } if(debug) { cout << "first stage ans = " << ans << "\n"; } vector<int> F(4); F[0] = 1; for(int it : son[id]) { int pick = 0, np = ksm(2, in[it]); int l = dfn[it], r = l + sz[it]; for(int p = l; p < r; p++) addt(pick, f[p]); vector<int> G(4); if(debug) { cout << "it = " << it << " pick = " << pick << " non pick = " << np << "\n"; } for(int i = 0; i < 4; i++) { addt(G[i], 1ll * F[i] * np % mod); addt(G[min(3, i + 1)], 1ll * F[i] * pick % mod); } F = G; } int pos = dfn[id]; for(int it : son[id]) { int l = dfn[it], r = l + sz[it]; int c = ksm(2, in[it]); for(int p = pos + 1; p < l; p++) f[p] = 1ll * f[p] * c % mod; for(int p = r; p < pos + sz[id]; p++) f[p] = 1ll * f[p] * c % mod; } if(vis[id]) f[pos] = mod - add(F[0], F[1]); else f[pos] = add(F[2], F[3]); addt(ans, 1ll * f[pos] * ksm(2, out[id]) % mod); addt(ans, mod - 1ll * F[2] * ksm(2, out[id]) % mod); // add this line if(debug) { cout << id << " pos = " << pos << " vis " << vis[id] << "\n"; for(int i = 1; i <= n; i++) cout << f[i] << " "; cout << "\n"; cout << "F: " << F[0] << " " << F[1] << " " << F[2] << " " << F[3] << "\n"; cout << "curans = " << ans << "\n"; } } void mian() { cin >> n >> m >> k; for(int i = 1, a, b; i < n; i++) { cin >> a >> b; e[a].push_back(b); e[b].push_back(a); } dfs1(1, 0); for(int i = 1; i <= __lg(n); i++) { for(int j = 1; j + (1 << i) - 1 <= n; j++) { mi[i][j] = get(mi[i - 1][j], mi[i - 1][j + (1 << i - 1)]); } } for(int i = 1, a, b; i <= m; i++) { cin >> a >> b; buc[a].push_back(b); buc[b].push_back(a); int d = lca(a, b); if(d != a && d != b) sub[d].push_back({a, b}); } for(int i = 1, a; i <= k; i++) cin >> a, vis[a] = 1; dfs2(1), dfs3(1), dfs4(1); if(debug) { cout << "ans = " << ans << "\n"; } cout << add(mod, -ans) << "\n"; } bool Med; int main() { // freopen("dfs0.in", "r", stdin); // freopen("1.out", "w", stdout); fprintf(stderr, "%.3lf MB\n", (&Mbe - &Med) / 1048576.0); ios::sync_with_stdio(0), cin.tie(0), cout.tie(0); int c, T = 1; cin >> c; while(T--) mian(); cerr << 1e3 * clock() / CLOCKS_PER_SEC << " ms\n"; return 0; }未卡常的正解
#include <bits/stdc++.h> using namespace std; using ll = long long; using pii = pair<int, int>; using pll = pair<ll, ll>; using pdi = pair<double, int>; using pdd = pair<double, double>; using ull = unsigned long long; #define ppc(x) __builtin_popcount(x) #define clz(x) __builtin_clz(x) bool Mbe; // mt19937 rnd(chrono::steady_clock::now().time_since_epoch().count()); mt19937_64 rnd(1064); int rd(int l, int r) { return rnd() % (r - l + 1) + l; } constexpr int mod = 1e9 + 7; void addt(int &x, int y) { x += y, x >= mod && (x -= mod); } int add(int x, int y) { return x += y, x >= mod && (x -= mod), x; } int ksm(int a, int b) { int s = 1; while(b) { if(b & 1) s = 1ll * s * a % mod; a = 1ll * a * a % mod, b >>= 1; } return s; } constexpr int Z = 1e6 + 5; int fc[Z], ifc[Z]; int bin(int n, int m) { if(n < m) return 0; return 1ll * fc[n] * ifc[m] % mod * ifc[n - m] % mod; } void init_fac(int Z) { for(int i = fc[0] = 1; i < Z; i++) fc[i] = 1ll * fc[i - 1] * i % mod; ifc[Z - 1] = ksm(fc[Z - 1], mod - 2); for(int i = Z - 2; ~i; i--) ifc[i] = 1ll * ifc[i + 1] * (i + 1) % mod; } void cmin(ll &x, ll y) { x = x < y ? x : y; } char buf[1 << 21], *p1 = buf, *p2 = buf; #define gc (p1 == p2 && (p2 = (p1 = buf) + fread(buf, 1, 1 << 21, stdin), p1 == p2) ? EOF : *p1++) inline int read() { int x = 0; char s = gc; while(!isdigit(s)) s = gc; while(isdigit(s)) x = x * 10 + s - '0', s = gc; return x; } // ---------- templates above ---------- constexpr int K = 20; constexpr int N = 5e5 + 5; int n, m, k, vis[N]; int pw[N], ipw[N], ans; vector<pii> sub[N]; vector<int> e[N], son[N], buc[N]; int dn, dfn[N], rev[N], sz[N], mi[K][N]; bool cmp(int x, int y) { return dfn[x] < dfn[y]; } int get(int x, int y) { return dfn[x] < dfn[y] ? x : y; } bool anc(int x, int y) { return dfn[x] <= dfn[y] && dfn[y] < dfn[x] + sz[x]; } int suc(int x, int y) { if(!anc(x, y)) return mi[0][dfn[x]]; int l = 0, r = son[x].size() - 1; while(l < r) { int m = l + r + 2 >> 1; if(dfn[y] >= dfn[son[x][m]]) l = m; else r = m - 1; } return son[x][l]; } int lca(int x, int y) { if(x == y) return x; if((x = dfn[x]) > (y = dfn[y])) swap(x, y); int d = __lg(y - x++); return get(mi[d][x], mi[d][y - (1 << d) + 1]); } void dfs1(int id, int ff) { sz[id] = 1; mi[0][dfn[id] = ++dn] = ff; rev[dn] = id; for(int it : e[id]) { if(it == ff) continue; dfs1(it, id); sz[id] += sz[it]; son[id].push_back(it); } sort(son[id].begin(), son[id].end(), cmp); } int in[N], out[N]; void dfs2(int id) { for(int it : son[id]) { dfs2(it), in[id] += in[it]; } for(int it : buc[id]) { in[id] += anc(id, it); } } void dfs3(int id) { int tot = 0; for(int it : buc[id]) { if(!anc(id, it)) out[id]++; else out[suc(id, it)]--, tot++; } int s = out[id]; for(int it : son[id]) s += in[it]; for(int it : son[id]) { out[it] += s - in[it] + tot, dfs3(it); } } int val[N << 2], laz[N << 2]; void tag(int x, int v) { laz[x] = 1ll * laz[x] * v % mod; val[x] = 1ll * val[x] * v % mod; } void down(int x) { if(laz[x] != 1) { tag(x << 1, laz[x]); tag(x << 1 | 1, laz[x]); laz[x] = 1; } } void modify(int l, int r, int p, int x, int v) { if(l == r) return val[x] = v, void(); int m = l + r >> 1; down(x); if(p <= m) modify(l, m, p, x << 1, v); else modify(m + 1, r, p, x << 1 | 1, v); val[x] = add(val[x << 1], val[x << 1 | 1]); } void modify(int l, int r, int ql, int qr, int x, int v) { if(ql > qr) return; if(ql <= l && r <= qr) return tag(x, v); int m = l + r >> 1; down(x); if(ql <= m) modify(l, m, ql, qr, x << 1, v); if(m < qr) modify(m + 1, r, ql, qr, x << 1 | 1, v); val[x] = add(val[x << 1], val[x << 1 | 1]); } int query(int l, int r, int ql, int qr, int x) { if(ql > qr) return 0; if(ql <= l && r <= qr) return val[x]; int m = l + r >> 1, ans = 0; down(x); if(ql <= m) ans = query(l, m, ql, qr, x << 1); if(m < qr) addt(ans, query(m + 1, r, ql, qr, x << 1 | 1)); return ans; } struct eve { int x, l, r, c; bool operator < (const eve &z) const { return x < z.x; } }; void dfs4(int id) { for(int it : son[id]) dfs4(it); for(int it : buc[id]) { if(!anc(id, it)) continue; int l = dfn[it], r = l + sz[it]; modify(1, n, l, r - 1, 1, 2); in[suc(id, it)]++; } int L = dfn[id], R = L + sz[id]; if(son[id].size() > 1) { vector<eve> arr; for(int it : son[id]) { int l = dfn[it], r = l + sz[it]; modify(1, n, l, r - 1, 1, ipw[in[it]]); // in -> _in } for(int it : son[id]) arr.push_back({dfn[it], 0, 0, 1}); // it -> dfn[it] for(pii I : sub[id]) { int u = I.first, v = I.second; if(dfn[u] > dfn[v]) swap(u, v); int l1 = dfn[u], r1 = l1 + sz[u]; int l2 = dfn[v], r2 = l2 + sz[v]; arr.push_back({l1, l2, r2 - 1, 2}); arr.push_back({r1, l2, r2 - 1, mod + 1 >> 1}); } sort(arr.begin(), arr.end()); auto find = [&](int x) { int l = 0, r = son[id].size() - 1; while(l < r) { int m = l + r >> 1; if(dfn[son[id][m]] > x) r = m; // l = m -> r = m, >= -> >, dfn[x] -> x else l = m + 1; } return dfn[son[id][l]]; }; int lst = L + 1; for(eve it : arr) { if(lst < it.x) { int c = query(1, n, lst, it.x - 1, 1); c = 1ll * c * query(1, n, find(lst), R - 1, 1) % mod; addt(ans, 1ll * c * pw[out[id] + in[id]] % mod); lst = it.x; } if(it.c != 1) { modify(1, n, it.l, it.r, 1, it.c); } } for(int it : son[id]) { int l = dfn[it], r = l + sz[it]; modify(1, n, l, r - 1, 1, pw[in[it]]); // in -> _in } } vector<int> F(4); F[0] = 1; for(int it : son[id]) { int l = dfn[it], r = l + sz[it]; int pick = query(1, n, l, r - 1, 1), np = pw[in[it]]; vector<int> G(4); for(int i = 0; i < 4; i++) { addt(G[i], 1ll * F[i] * np % mod); addt(G[min(3, i + 1)], 1ll * F[i] * pick % mod); } F = G; } for(int it : son[id]) { int l = dfn[it], r = l + sz[it], c = pw[in[it]]; modify(1, n, L + 1, l - 1, 1, c); modify(1, n, r, R - 1, 1, c); } int f = vis[id] ? mod - add(F[0], F[1]) : add(F[2], F[3]); modify(1, n, L, 1, f); addt(ans, 1ll * add(f, mod - F[2]) * pw[out[id]] % mod); } void mian() { n = read(), m = read(), k = read(); pw[0] = ipw[0] = 1; for(int i = 1; i <= m; i++) { pw[i] = add(pw[i - 1], pw[i - 1]); ipw[i] = 1ll * ipw[i - 1] * (mod + 1 >> 1) % mod; } for(int i = 1, a, b; i < n; i++) { a = read(), b = read(); e[a].push_back(b); e[b].push_back(a); } dfs1(1, 0); for(int i = 1; i <= __lg(n); i++) { for(int j = 1; j + (1 << i) - 1 <= n; j++) { mi[i][j] = get(mi[i - 1][j], mi[i - 1][j + (1 << i - 1)]); } } for(int i = 1, a, b; i <= m; i++) { a = read(), b = read(); buc[a].push_back(b); buc[b].push_back(a); int d = lca(a, b); if(d != a && d != b) sub[d].push_back({a, b}); } for(int i = 1; i <= k; i++) vis[read()] = 1; dfs2(1), dfs3(1), dfs4(1); cout << add(mod, -ans) << "\n"; } bool Med; int main() { fprintf(stderr, "%.3lf MB\n", (&Mbe - &Med) / 1048576.0); int c, T = 1; cin >> c; while(T--) mian(); cerr << 1e3 * clock() / CLOCKS_PER_SEC << " ms\n"; return 0; }卡常后的正解
#include <bits/stdc++.h> using namespace std; using ll = long long; using pii = pair<int, int>; using pll = pair<ll, ll>; using pdi = pair<double, int>; using pdd = pair<double, double>; using ull = unsigned long long; #define ppc(x) __builtin_popcount(x) #define clz(x) __builtin_clz(x) bool Mbe; // mt19937 rnd(chrono::steady_clock::now().time_since_epoch().count()); mt19937_64 rnd(1064); int rd(int l, int r) { return rnd() % (r - l + 1) + l; } constexpr int mod = 1e9 + 7; void addt(int &x, int y) { x += y, x >= mod && (x -= mod); } int add(int x, int y) { return x += y, x >= mod && (x -= mod), x; } int ksm(int a, int b) { int s = 1; while(b) { if(b & 1) s = 1ll * s * a % mod; a = 1ll * a * a % mod, b >>= 1; } return s; } constexpr int Z = 1e6 + 5; int fc[Z], ifc[Z]; int bin(int n, int m) { if(n < m) return 0; return 1ll * fc[n] * ifc[m] % mod * ifc[n - m] % mod; } void init_fac(int Z) { for(int i = fc[0] = 1; i < Z; i++) fc[i] = 1ll * fc[i - 1] * i % mod; ifc[Z - 1] = ksm(fc[Z - 1], mod - 2); for(int i = Z - 2; ~i; i--) ifc[i] = 1ll * ifc[i + 1] * (i + 1) % mod; } void cmin(ll &x, ll y) { x = x < y ? x : y; } char buf[1 << 21], *p1 = buf, *p2 = buf; #define gc (p1 == p2 && (p2 = (p1 = buf) + fread(buf, 1, 1 << 21, stdin), p1 == p2) ? EOF : *p1++) inline int read() { int x = 0; char s = gc; while(!isdigit(s)) s = gc; while(isdigit(s)) x = x * 10 + s - '0', s = gc; return x; } // ---------- templates above ---------- constexpr int K = 20; constexpr int N = 5e5 + 5; struct Linklist { int cnt, hd[N], nxt[N << 1], to[N << 1]; void add(int u, int v) { nxt[++cnt] = hd[u], hd[u] = cnt, to[cnt] = v; } } e, buc; struct Linklist_2 { int cnt, hd[N], nxt[N]; pii to[N]; void add(int u, pii v) { nxt[++cnt] = hd[u], hd[u] = cnt, to[cnt] = v; } } sub; int n, m, k, vis[N]; int pw[N], ipw[N], ans; vector<int> son[N]; int dn, dfn[N], rev[N], sz[N], mi[K][N]; bool cmp(int x, int y) { return dfn[x] < dfn[y]; } int get(int x, int y) { return dfn[x] < dfn[y] ? x : y; } bool anc(int x, int y) { return dfn[x] <= dfn[y] && dfn[y] < dfn[x] + sz[x]; } int suc(int x, int y) { if(!anc(x, y)) return mi[0][dfn[x]]; int l = 0, r = son[x].size() - 1; while(l < r) { int m = l + r + 2 >> 1; if(dfn[y] >= dfn[son[x][m]]) l = m; else r = m - 1; } return son[x][l]; } int lca(int x, int y) { if(x == y) return x; if((x = dfn[x]) > (y = dfn[y])) swap(x, y); int d = __lg(y - x++); return get(mi[d][x], mi[d][y - (1 << d) + 1]); } void dfs1(int id, int ff) { sz[id] = 1; mi[0][dfn[id] = ++dn] = ff; rev[dn] = id; for(int _ = e.hd[id]; _; _ = e.nxt[_]) { int it = e.to[_]; if(it == ff) continue; dfs1(it, id); sz[id] += sz[it]; son[id].push_back(it); } sort(son[id].begin(), son[id].end(), cmp); } int in[N], out[N]; void dfs2(int id) { for(int it : son[id]) { dfs2(it), in[id] += in[it]; } for(int _ = buc.hd[id]; _; _ = buc.nxt[_]) { int it = buc.to[_]; in[id] += anc(id, it); } } void dfs3(int id) { int tot = 0; for(int _ = buc.hd[id]; _; _ = buc.nxt[_]) { int it = buc.to[_]; if(!anc(id, it)) out[id]++; else out[suc(id, it)]--, tot++; } int s = out[id]; for(int it : son[id]) s += in[it]; for(int it : son[id]) { out[it] += s - in[it] + tot, dfs3(it); } } int val[N << 2], laz[N << 2]; void tag(int x, int v) { laz[x] = 1ll * laz[x] * v % mod; val[x] = 1ll * val[x] * v % mod; } void down(int x) { if(laz[x] != 1) { tag(x << 1, laz[x]); tag(x << 1 | 1, laz[x]); laz[x] = 1; } } void modify(int l, int r, int p, int x, int v) { if(l == r) return val[x] = v, void(); int m = l + r >> 1; down(x); if(p <= m) modify(l, m, p, x << 1, v); else modify(m + 1, r, p, x << 1 | 1, v); val[x] = add(val[x << 1], val[x << 1 | 1]); } void modify(int l, int r, int ql, int qr, int x, int v) { if(ql > qr) return; if(ql <= l && r <= qr) return tag(x, v); int m = l + r >> 1; down(x); if(ql <= m) modify(l, m, ql, qr, x << 1, v); if(m < qr) modify(m + 1, r, ql, qr, x << 1 | 1, v); val[x] = add(val[x << 1], val[x << 1 | 1]); } int query(int l, int r, int ql, int qr, int x) { if(ql > qr) return 0; if(ql <= l && r <= qr) return val[x]; int m = l + r >> 1, ans = 0; down(x); if(ql <= m) ans = query(l, m, ql, qr, x << 1); if(m < qr) addt(ans, query(m + 1, r, ql, qr, x << 1 | 1)); return ans; } struct eve { int x, l, r, c; bool operator < (const eve &z) const { return x < z.x; } }; void dfs4(int id) { for(int it : son[id]) dfs4(it); for(int _ = buc.hd[id]; _; _ = buc.nxt[_]) { int it = buc.to[_]; if(!anc(id, it)) continue; int l = dfn[it], r = l + sz[it]; modify(1, n, l, r - 1, 1, 2); in[suc(id, it)]++; } int L = dfn[id], R = L + sz[id]; vector<int> F(4); F[0] = 1; for(int it : son[id]) { int l = dfn[it], r = l + sz[it]; int pick = query(1, n, l, r - 1, 1), np = pw[in[it]]; vector<int> G(4); for(int i = 0; i < 4; i++) { addt(G[i], 1ll * F[i] * np % mod); addt(G[min(3, i + 1)], 1ll * F[i] * pick % mod); } F = G; } for(int it : son[id]) { int l = dfn[it], r = l + sz[it]; modify(1, n, l, r - 1, 1, ipw[in[it]]); } if(son[id].size() > 1) { vector<eve> arr; for(int it : son[id]) arr.push_back({dfn[it], 0, 0, 1}); // it -> dfn[it] for(int _ = sub.hd[id]; _; _ = sub.nxt[_]) { pii it = sub.to[_]; int u = it.first, v = it.second; if(dfn[u] > dfn[v]) swap(u, v); int l1 = dfn[u], r1 = l1 + sz[u]; int l2 = dfn[v], r2 = l2 + sz[v]; arr.push_back({l1, l2, r2 - 1, 2}); arr.push_back({r1, l2, r2 - 1, mod + 1 >> 1}); } sort(arr.begin(), arr.end()); auto find = [&](int x) { int l = 0, r = son[id].size() - 1; while(l < r) { int m = l + r >> 1; if(dfn[son[id][m]] > x) r = m; // l = m -> r = m, >= -> >, dfn[x] -> x else l = m + 1; } return dfn[son[id][l]]; }; int lst = L + 1; for(eve it : arr) { if(lst < it.x) { int c = query(1, n, lst, it.x - 1, 1); c = 1ll * c * query(1, n, find(lst), R - 1, 1) % mod; addt(ans, 1ll * c * pw[out[id] + in[id]] % mod); lst = it.x; } if(it.c != 1) { modify(1, n, it.l, it.r, 1, it.c); } } } modify(1, n, L + 1, R - 1, 1, pw[in[id]]); int f = vis[id] ? mod - add(F[0], F[1]) : add(F[2], F[3]); modify(1, n, L, 1, f); addt(ans, 1ll * add(f, mod - F[2]) * pw[out[id]] % mod); } void mian() { n = read(), m = read(), k = read(); pw[0] = ipw[0] = 1; for(int i = 1; i <= m; i++) { pw[i] = add(pw[i - 1], pw[i - 1]); ipw[i] = 1ll * ipw[i - 1] * (mod + 1 >> 1) % mod; } for(int i = 1, a, b; i < n; i++) { a = read(), b = read(); e.add(a, b), e.add(b, a); } dfs1(1, 0); for(int i = 1; i <= __lg(n); i++) { for(int j = 1; j + (1 << i) - 1 <= n; j++) { mi[i][j] = get(mi[i - 1][j], mi[i - 1][j + (1 << i - 1)]); } } for(int i = 1, a, b; i <= m; i++) { a = read(), b = read(); buc.add(a, b); buc.add(b, a); int d = lca(a, b); if(d != a && d != b) sub.add(d, {a, b}); } for(int i = 1; i <= k; i++) vis[read()] = 1; dfs2(1), dfs3(1), dfs4(1); cout << add(mod, -ans) << "\n"; } bool Med; int main() { fprintf(stderr, "%.3lf MB\n", (&Mbe - &Med) / 1048576.0); int c, T = 1; cin >> c; while(T--) mian(); cerr << 1e3 * clock() / CLOCKS_PER_SEC << " ms\n"; return 0; }
- 1
信息
- ID
- 7291
- 时间
- 2000ms
- 内存
- 512MiB
- 难度
- 10
- 标签
- 递交数
- 8
- 已通过
- 1
- 上传者