求助,为什么会 T /kk
查看原帖
求助,为什么会 T /kk
183881
_shy楼主2023/3/22 23:49
#include <bits/stdc++.h>
#define ll long long
#define ld long double
#define ull unsigned int
#define lson i << 1
#define rson i << 1 | 1
using namespace std;
const int maxn = 5e5 + 10;
int a[maxn], n, m, q;
// 快读板子
void read (int &x) 
{
	int w = 1; x = 0;
	char ch = getchar ();
	while (ch < '0' || ch > '9') 
	{
		if (ch == '-') w = -1;
		ch = getchar (); 
	}
	while (ch >= '0' && ch <= '9') 
	{
		x = (x << 1) + (x << 3) + ch - 48;
		ch = getchar ();
	}
	x *= w;
} 
// 预处理 base
ull base[2][maxn];
void Init_Base () 
{
	int len = max (n, m);
	base[0][0] = base[1][0] = 1;
	for (int o = 0; o <= 1; o++)
		for (int i = 1; i <= len + 5; i++)
			base[o][i] = base[o][i - 1] * (ull) (o == 0 ? 1000003 : 1000033);
} 
// 建立 Segment_Tree 存 a 的 哈希值
int st[maxn], tp; // 存 [l,r] 有关区间 
struct Segment_Tree 
{
	struct arr 
	{
		int l, r;
		ull sum[2];
	} t[maxn << 2];
	void build (int i, int l, int r) 
	{
		t[i].l = l, t[i].r = r;
		if (l == r)
		{
			t[i].sum[0] = t[i].sum[1] = a[l];
			return;
		}
		int mid = (l + r) >> 1;
		build (lson, l, mid), build (rson, mid + 1, r);
		for (int o = 0; o <= 1; o++)
			t[i].sum[o] = t[lson].sum[o] + t[rson].sum[o] * base[o][mid - l + 1]; 
	}
	void merge (int i, int pos, ull k) 
	{
		if (t[i].l == t[i].r && t[i].l == pos) 
		{
			t[i].sum[0] = t[i].sum[1] = k;
			return;
		}
		if (pos <= t[lson].r) merge (lson, pos, k);
		else merge (rson, pos, k);
		int mid = (t[i].l + t[i].r) >> 1;
		for (int o = 0; o <= 1; o++)
			t[i].sum[o] = t[lson].sum[o] + t[rson].sum[o] * base[o][(mid - t[i].l + 1)];
	}
	void findst (int i, int l, int r) 
	{
		if (t[i].r < l || t[i].l > r) return;
		if (t[i].l >= l && t[i].r <= r) 
		{
			st[++ tp] = i;
			return;
		}
		findst (lson, l, r), findst (rson, l, r);
	} 
} Hash_Arr; 
// 构建树上 Hash 
map<pair<ull, ull>, int> mp;
vector<int> son[maxn];
int depth[maxn], rt; 
ull Hash[2][maxn];
void dfs (int u, int d) 
{
	sort (son[u].begin (), son[u].end ());
	vector<int> :: iterator it = son[u].begin ();
	int p = 1;
	for (it; it != son[u].end (); it ++) 
	{
		int v = *it; 
		depth[v] = d + 1;
		for (int o = 0; o <= 1; o++)
			Hash[o][v] = Hash[o][u] + base[o][d] * p; 
		mp.insert (make_pair (make_pair (Hash[0][v], Hash[1][v]), v));
		p ++, dfs (v, d + 1);
	}
	return;
}
void ask (int x, int l, int r) 
{
	tp = 0, Hash_Arr.findst (1, l, r);
	ull cur[2]; int ans = x;
	for (int o = 0; o <= 1; o++) cur[o] = Hash[o][x];
	for (int i = 1; i <= tp; i++) 
	{
		ull curi[2]; int y = st[i];
		for (int o = 0; o <= 1; o++) 
			curi[o] = cur[o] + Hash_Arr.t[y].sum[o] * base[o][Hash_Arr.t[y].l - l] * base[o][depth[x]];
		if (mp.find (make_pair (curi[0], curi[1])) != mp.end ()) 
		{
			ans = mp[make_pair (curi[0], curi[1])];
			for (int o = 0; o <= 1; o++) cur[o] = curi[o];
		}
		else 
		{
			y = st[i];
			int li = Hash_Arr.t[y].l, ri = Hash_Arr.t[y].r;
			while (li < ri) 
			{
				for (int o = 0; o <= 1; o++) 
				curi[o] = cur[o] + Hash_Arr.t[y << 1].sum[o] * base[o][Hash_Arr.t[y << 1].l - l] * base[o][depth[x]];
				if (mp.find (make_pair (curi[0], curi[1])) != mp.end ()) 
				{
					ans = mp[make_pair (curi[0], curi[1])];
					li = Hash_Arr.t[y << 1].r + 1, y = y << 1 | 1;
					for (int o = 0; o <= 1; o++) cur[o] = curi[o];
				}
				else 
					ri = Hash_Arr.t[y << 1].r, y = y << 1;
			}
			break;
		}
	}
	printf ("%d\n", ans);
}
int main ()
{
	read (n), read (m), read (q);
	Init_Base ();
	for (int i = 1; i <= n; i++) 
	{
		int fa; read (fa);
		if (!fa) rt = i;
		son[fa].push_back (i);
	}
	Hash[0][rt] = Hash[1][rt] = 1, depth[rt] = 1;
	dfs (rt, 1);
	for (int i = 1; i <= m; i++)
		read (a[i]);
	Hash_Arr.build (1, 1, m);
	while (q --) 
	{
		int opt, l, r, x;
		read (opt);
		if (opt == 1) 
		{
			read (x), read (l), read (r);
			ask (x, l, r);
		}
		else 
		{
			read (l), read (x);
			Hash_Arr.merge (1, l, x);
		}
	}
	return 0;
}


2023/3/22 23:49
加载中...