Not a member of Pastebin yet?
Sign Up,
it unlocks many cool features!
- //jb workshop
- //misaka and elaina will carry me to master
- #include <iostream>
- #include <cstdio>
- #include <cstring>
- #include <cmath>
- #include <utility>
- #include <cassert>
- #include <algorithm>
- #include <vector>
- #include <functional>
- #include <numeric>
- #include <set>
- #include <array>
- #include <queue>
- #include <map>
- #include <chrono>
- #include <random>
- #define ll long long
- #define lb long double
- #define sz(vec) ((int)(vec.size()))
- #define all(x) x.begin(), x.end()
- #define pb push_back
- #define mp make_pair
- #define kill(x, s) {int COND = x; if(COND){ cout << s << "\n"; return ; }}
- #ifdef ONLINE_JUDGE
- #define cerr while(0) cerr
- #endif
- const lb eps = 1e-9;
- const ll mod = 1e9 + 7, ll_max = 1e18;
- //const ll mod = (1 << (23)) * 119 +1, ll_max = 1e18;
- const int MX = 2e5 +10, int_max = 0x3f3f3f3f;
- struct {
- template<class T>
- operator T() {
- T x; std::cin >> x; return x;
- }
- } in;
- using namespace std;
- struct hld{
- vector<int> dep, sz, par, head, tin, out;
- vector<vector<int>> adj;
- int n, ind;
- hld(){}
- hld(vector<pair<int, int>> edges){
- n = sz(edges) + 1;
- //assert(n == 5);
- ind = 0;
- dep = sz = par = head = tin = out = vector<int>(n+1, 0);
- adj = vector<vector<int>>(n+1);
- for(auto [a, b] : edges){
- adj[a].pb(b);
- adj[b].pb(a);
- }
- dfs(1, 0);
- head[1] = 1;
- dfs2(1, 0);
- cerr << "pr tree\n";
- for(int i = 1; i<=n; i++){
- cerr << i << " " << par[i] << " " << head[i] << "\n";
- }
- }
- void dfs(int x, int p){
- sz[x] = 1;
- dep[x] = dep[p] + 1;
- par[x] = p;
- for(auto &i : adj[x]){
- if(i == p) continue;
- dfs(i, x);
- sz[x] += sz[i];
- if(adj[x][0] == p || sz[i] > sz[adj[x][0]]) swap(adj[x][0], i);
- }
- }
- void dfs2(int x, int p){
- tin[x] = ind++;
- for(auto &i : adj[x]){
- if(i == p) continue;
- head[i] = (i == adj[x][0] ? head[x] : i);
- dfs2(i, x);
- }
- out[x] = ind;
- }
- int lca(int a, int b){
- while(head[a] != head[b]){
- if(dep[head[a]] > dep[head[b]]) swap(a, b);
- b = par[head[b]];
- }
- if(dep[a] > dep[b]) swap(a, b);
- return a;
- }
- };
- const int s = (1 << 18);
- hld A, B;
- vector<pair<int, int>> edges1, edges2;
- vector<int> tour;
- int n, q, P;
- int root[MX];
- const int NN = MX * 60;
- int lc[NN], rc[NN], sum[NN];
- int ind = 0;
- bool cmp(int a, int b){
- return A.tin[a] < A.tin[b];
- }
- void dup(int& k){
- ind++;
- lc[ind] = lc[k];
- rc[ind] = rc[k];
- sum[ind] = sum[k];
- k = ind;
- }
- void U(int p, int& k, int L, int R){
- dup(k);
- if(L + 1 == R){
- sum[k]++;
- return ;
- }
- int mid = (L + R)/2;
- if(p < mid) U(p, lc[k], L, mid);
- else U(p, rc[k], mid, R);
- sum[k] = sum[lc[k]] + sum[rc[k]];
- }
- int S(int qL, int qR, int k, int L, int R){
- if(qR <= L || R <= qL) return 0;
- if(qL <= L && R <= qR){
- return sum[k];
- }
- int mid = (L + R)/2;
- return S(qL, qR, lc[k], L, mid) + S(qL, qR, rc[k], mid, R);
- }
- int query(int a, int u, int v){
- int ans = 0;
- //cerr << u << " " << v << "\n";
- //int cnt = 0;
- while(B.head[u] != B.head[v]){
- //cnt++;
- //if(cnt > 10) break;
- //cerr << u << " " << v << "\n";
- if(B.dep[B.head[u]] < B.dep[B.head[v]]) swap(u, v);
- ans += S(B.tin[B.head[u]], B.tin[u]+1, root[a], 0, s);
- u = B.par[B.head[u]];
- }
- if(B.dep[u] > B.dep[v]) swap(u, v);
- ans += S(B.tin[u], B.tin[v]+1, root[a], 0, s);
- return ans;
- }
- #define res(a) ((a + last*P - 1)%n +1)
- void solve(){
- n = in, q = in, P = in;
- for(int i = 1; i<n; i++){
- int a =in, b = in;
- edges1.pb(mp(a, b));
- }
- for(int i = 1; i<n; i++){
- int a =in, b= in;
- edges2.pb(mp(a, b));
- }
- cerr << "\n";
- A = hld(edges1);
- B = hld(edges2);
- tour = vector<int>(n);
- iota(all(tour), 1);
- sort(all(tour), cmp);
- cerr << "\n";
- for(int x : tour){
- root[x] = root[A.par[x]];
- cerr << x << " ";
- U(B.tin[x], root[x], 0, s);
- }
- cerr << "\n";
- int last = 0;
- for(int qaq = 1; qaq <= q; qaq++){
- int a = in, b = in, c = in, d = in;
- a = res(a);
- b = res(b);
- c = res(c);
- d = res(d);
- cerr << a << " " << b << " " << c << " " << d << "\n";
- int l = A.lca(a, b), p = A.par[l];
- int a1 = query(a, c, d);
- int a2 = query(b, c, d);
- int a3 = query(l, c, d);
- int a4 = query(p, c, d);
- cerr << a << " " << a1 << "\n";
- cerr << b << " " << a2 << "\n";
- cerr << l << " " << a3 << "\n";
- cerr << p << " " << a4 << "\n";
- last = a1 + a2 - a3 - a4;
- cout << last << "\n";
- }
- }
- signed main(){
- cin.tie(0) -> sync_with_stdio(0);
- int T = 1;
- //cin >> T;
- for(int i = 1; i<=T; i++){
- solve();
- }
- return 0;
- }
Advertisement
Add Comment
Please, Sign In to add comment