[虚树模板] 求助一个小问题(已A)
查看原帖
[虚树模板] 求助一个小问题(已A)
762588
Edgebright楼主2023/8/21 22:11

为什么N开两倍才能过,开2.5e5不够?

#include<bits/stdc++.h>
#define int long long
#define __lg(x) ((x)? __lg(x) : 0)
using namespace std;
const int N = 500505, L = 21, I = 0x3f3f3f3f3f3f3f3f;//N开250505 会WA on #10
int n;
struct edge
{
	int n, t, w;
}e[N << 1];
int h[N], ce = 1;
inline void add(int u, int v, int w)
{
	e[++ce] = {h[u], v, w}; h[u] = ce;
	return;
}

int dfn[N], stp, minw[N];
int mn[L][N];
inline int mndfn(int x, int y)
{
	return (dfn[x] < dfn[y])? x : y;
}
void dfs(int u, int f)
{
	mn[0][dfn[u] = ++stp] = f;
	for(int i = h[u]; i; i = e[i].n)
	{
		int to = e[i].t;
		if(to == f) continue;
		minw[to] = min(minw[u], e[i].w);
		dfs(to, u);
	}
	return;
}
void SparseTable()
{
	int ex = __lg(n);
	for(int i = 1; i <= ex; ++i)
		for(int j = 1; j <= n; ++j)
			mn[i][j] = mndfn(mn[i - 1][j], mn[i - 1][j + (1 << (i - 1))]);
	return;
}
inline int lca(int u, int v)
{
	if(u == v) return u;
	u = dfn[u]; v = dfn[v]; if(u > v) swap(u, v);
	++u; int d = __lg(v - u);
	return mndfn(mn[d][u], mn[d][v + 1 - (1 << d)]);
}

int m;
int c[N], rich[N];
int id;
int dp(int u)
{
//	printf("u%lld\n", u);
	if(rich[u] == id) return I;
	int res = 0;
	for(int i = h[u]; i; i = e[i].n)
	{
		res += min(e[i].w, dp(e[i].t));
	}
	
	return res;
}
signed main()
{
	scanf("%lld", &n);
	memset(minw, 0x3f, (n + 10) * 8);
	for(int i = 1; i < n; ++i)
	{
		int u, v, w;
		scanf("%lld%lld%lld", &u, &v, &w);
		add(u, v, w); add(v, u, w);
	}
	dfs(1, 0);
	SparseTable(); 
	scanf("%lld", &m);
	for(int i = 1; i <= m; ++i)
	{
		int k;
		scanf("%lld", &k);
		ce = 1; id = i;
		c[1] = 1; h[1] = 0;
		for(int j = 2; j <= k + 1; ++j)
		{
			scanf("%lld", &c[j]);
			h[c[j]] = 0; rich[c[j]] = id;
		}
		sort(c + 1, c + k + 2, [](int x, int y){return dfn[x] < dfn[y];});
		int len = k + 1;
		for(int j = 1; j <= k; ++j)
		{
			c[++len] = lca(c[j], c[j + 1]);
			h[c[len]] = 0;
		}	
		sort(c + 1, c + len + 1, [](int x, int y){return dfn[x] < dfn[y];});
		len = unique(c + 1, c + len + 1) - c - 1;
//		for(int j = 1; j <= len; ++j)
//		{
//			printf("%lld ", c[j]);
//		}puts("");
		for(int j = 1; j < len; ++j)
		{
			int lc = lca(c[j], c[j + 1]);
			add(lc, c[j + 1], minw[c[j + 1]]);
		}
		printf("%lld\n", dp(1));
	}
	return 0;
}
2023/8/21 22:11
加载中...