求助,只 AC 了第 8 个点
查看原帖
求助,只 AC 了第 8 个点
183562
_ReClouds_楼主2022/5/10 19:19

基本思路是题解的第一种做法。

代码略去缺省源和头文件

namespace Program
{
	#define BL(i) (s * (i - 1) + 1)
	#define BR(i) (min(len, s * i))

    const int MAXN = 40005;
	const int MAXM = 305;

    int n, m, s, col[MAXN], tot = 1, hd[MAXN], to[MAXN << 1], nxt[MAXN << 1], len, seq[MAXN << 1], pos[MAXN][2], in[MAXN << 1], cnt[2][MAXM][MAXN], st[MAXM][MAXM], tmp[2][MAXN], fa[MAXN], dep[MAXN], si[MAXN], hs[MAXN], top[MAXN], ans;

	inline void Link(int u, int v) { return ++tot, nxt[tot] = hd[u], to[tot] = v, hd[u] = tot, void(); }

    inline void Discretization()
    {
        sort(seq + 1, seq + len + 1), len = unique(seq + 1, seq + len + 1) - seq - 1;
        for(register int i = 1; i <= n; i++) col[i] = lower_bound(seq + 1, seq + len + 1, col[i]) - seq;
        return len = 0, void();
    }

	inline int LCA(int u, int v)
	{
		while(top[u] != top[v])
		{
			if(dep[top[u]] < dep[top[v]]) swap(u, v);
			u = fa[top[u]];
		}
		if(dep[u] < dep[v]) return u;
		return v;
	}

	inline void DFS1(int u, int f)
	{
		seq[++len] = u, pos[u][0] = len, fa[u] = f, dep[u] = dep[f] + 1, si[u] = 1, hs[u] = 0;
		for(register int i = hd[u]; i; i = nxt[i])
		{
			int v = to[i];
			if(v == f) continue;
			DFS1(v, u);
			si[u] += si[v];
			if(si[v] > si[hs[u]]) hs[u] = v;
		}
		return seq[++len] = u, pos[u][1] = len, void();
	}

	inline void DFS2(int u, int t)
	{
		top[u] = t;
		if(hs[u]) DFS2(hs[u], t);
		for(register int i = hd[u]; i; i = nxt[i])
		{
			int v = to[i];
			if(v == fa[u] || v == hs[u]) continue;
			DFS2(v, v);
		}
		return;
	}

	inline void Extend(int u, int *tmp0, int *tmp1, int &res)
	{
		++tmp0[u];
		if(tmp0[u] == 1) res += (!tmp1[col[u]]), ++tmp1[col[u]];
		else --tmp1[col[u]], res -= (!tmp1[col[u]]);
		return;
	}

	inline void Init()
	{
		s = (int)sqrt(len);
		for(register int i = 1; i <= len; i++)
		{
			in[i] = (i - 1) / s + 1, ++tmp[0][seq[i]], ++tmp[1][col[seq[i]]];
			if(i == BR(in[i])) for(register int j = 1; j <= n; j++) cnt[0][in[i]][j] = tmp[0][j], cnt[1][in[i]][j] = tmp[1][j];
		}
		memset(tmp, 0, sizeof tmp);
		for(register int i = 1; i <= in[len]; i++)
		{
			int res = 0;
			memset(tmp, 0, sizeof tmp);
			for(register int j = BL(i); j <= len; j++)
			{
				Extend(seq[j], tmp[0], tmp[1], res);
				if(j == BR(in[j])) st[i][in[j]] = res;
			}
		}
		return memset(tmp, 0, sizeof tmp), void();
	}

    inline i32 Run()
    {
        n = Read32(), m = Read32();
        for(register int i = 1; i <= n; i++) seq[++len] = col[i] = Read32();
        Discretization();
        for(register int i = 1; i < n; i++)
        {
			int u = Read32(), v = Read32();
			Link(u, v), Link(v, u);
        }
		DFS1(1, 0), DFS2(1, 1), Init();
		while(m--)
		{
			int u = Read32() ^ ans, v = Read32(), t = LCA(u, v);
			if(u != t && v != t)
			{
				if(pos[u][1] > pos[v][0]) swap(u, v);
				int pu = pos[u][1], pv = pos[v][0];
				if(in[pu] == in[pv])
				{
					ans = 0;
					Extend(t, tmp[0], tmp[1], ans);
					for(register int i = pu; i <= pv; i++) Extend(seq[i], tmp[0], tmp[1], ans);
					Write32(ans);
					tmp[0][t] = tmp[1][col[t]] = 0;
					for(register int i = pu; i <= pv; i++) tmp[0][seq[i]] = tmp[1][col[seq[i]]] = 0;
					continue;
				}
				int lb = in[pu] + 1, rb = in[pv] - 1;
				ans = st[lb][rb];
				tmp[0][t] = cnt[0][rb][t] - cnt[0][lb - 1][t], tmp[1][col[t]] = cnt[1][rb][col[t]] - cnt[1][lb - 1][col[t]];
				for(register int i = pu; i <= BR(in[pu]); i++) tmp[0][seq[i]] = cnt[0][rb][seq[i]] - cnt[0][lb - 1][seq[i]], tmp[1][col[seq[i]]] = cnt[1][rb][col[seq[i]]] - cnt[1][lb - 1][col[seq[i]]];
				for(register int i = BL(in[pv]); i <= pv; i++) tmp[0][seq[i]] = cnt[0][rb][seq[i]] - cnt[0][lb - 1][seq[i]], tmp[1][col[seq[i]]] = cnt[1][rb][col[seq[i]]] - cnt[1][lb - 1][col[seq[i]]];
				Extend(t, tmp[0], tmp[1], ans);
				for(register int i = pu; i <= BR(in[pu]); i++) Extend(seq[i], tmp[0], tmp[1], ans);
				for(register int i = BL(in[pv]); i <= pv; i++) Extend(seq[i], tmp[0], tmp[1], ans);
				Write32(ans);
				tmp[0][t] = tmp[1][col[t]] = 0;
				for(register int i = pu; i <= BR(in[pu]); i++) tmp[0][seq[i]] = tmp[1][col[seq[i]]] = 0;
				for(register int i = BL(in[pv]); i <= pv; i++) tmp[0][seq[i]] = tmp[1][col[seq[i]]] = 0;
			}
			else
			{
				if(pos[u][0] > pos[v][0]) swap(u, v);
				int pu = pos[u][0], pv = pos[v][0];
				if(in[pu] == in[pv])
				{
					ans = 0;
					for(register int i = pu; i <= pv; i++) Extend(seq[i], tmp[0], tmp[1], ans);
					Write32(ans);
					for(register int i = pu; i <= pv; i++) tmp[0][seq[i]] = tmp[1][col[seq[i]]] = 0;
					continue;
				}
				int lb = in[pu] + 1, rb = in[pv] - 1;
				ans = st[lb][rb];
				for(register int i = pu; i <= BR(in[pu]); i++) tmp[0][seq[i]] = cnt[0][rb][seq[i]] - cnt[0][lb - 1][seq[i]], tmp[1][col[seq[i]]] = cnt[1][rb][col[seq[i]]] - cnt[1][lb - 1][col[seq[i]]];
				for(register int i = BL(in[pv]); i <= pv; i++) tmp[0][seq[i]] = cnt[0][rb][seq[i]] - cnt[0][lb - 1][seq[i]], tmp[1][col[seq[i]]] = cnt[1][rb][col[seq[i]]] - cnt[1][lb - 1][col[seq[i]]];
				for(register int i = pu; i <= BR(in[pu]); i++) Extend(seq[i], tmp[0], tmp[1], ans);
				for(register int i = BL(in[pv]); i <= pv; i++) Extend(seq[i], tmp[0], tmp[1], ans);
				Write32(ans);
				for(register int i = pu; i <= BR(in[pu]); i++) tmp[0][seq[i]] = tmp[1][col[seq[i]]] = 0;
				for(register int i = BL(in[pv]); i <= pv; i++) tmp[0][seq[i]] = tmp[1][col[seq[i]]] = 0;
			}
		}
        return 0;
    }
	
	#undef BL
	#undef BR
}

i32 main() { return Program :: Run(); }
2022/5/10 19:19
加载中...