willy108

stupid hld

Jul 23rd, 2022
1,114
0
Never
Not a member of Pastebin yet? Sign Up, it unlocks many cool features!
C++ 4.48 KB | None | 0 0
  1.  
  2.  
  3. //jb workshop
  4.  
  5. //misaka and elaina will carry me to master
  6. #include <iostream>
  7. #include <cstdio>
  8. #include <cstring>
  9. #include <cmath>
  10. #include <utility>
  11. #include <cassert>
  12. #include <algorithm>
  13. #include <vector>
  14. #include <functional>
  15. #include <numeric>
  16. #include <set>
  17. #include <array>
  18. #include <queue>
  19. #include <map>
  20. #include <chrono>
  21. #include <random>
  22.  
  23. #define ll long long
  24. #define lb long double
  25. #define sz(vec) ((int)(vec.size()))
  26. #define all(x) x.begin(), x.end()
  27. #define pb push_back
  28. #define mp make_pair
  29. #define kill(x, s) {int COND = x; if(COND){ cout << s << "\n"; return ; }}
  30.  
  31. #ifdef ONLINE_JUDGE
  32. #define cerr while(0) cerr
  33. #endif
  34.  
  35. const lb eps = 1e-9;
  36. const ll mod = 1e9 + 7, ll_max = 1e18;
  37. //const ll mod = (1 << (23)) * 119 +1, ll_max = 1e18;
  38. const int MX = 2e5 +10, int_max = 0x3f3f3f3f;
  39.  
  40. struct {
  41.   template<class T>
  42.   operator T() {
  43.     T x; std::cin >> x; return x;
  44.   }
  45. } in;
  46.  
  47. using namespace std;
  48.  
  49. struct hld{
  50.     vector<int> dep, sz, par, head, tin, out;
  51.     vector<vector<int>> adj;
  52.     int n, ind;
  53.     hld(){}
  54.     hld(vector<pair<int, int>> edges){
  55.         n = sz(edges) + 1;
  56.         //assert(n == 5);
  57.         ind = 0;
  58.         dep = sz = par = head = tin = out = vector<int>(n+1, 0);
  59.         adj = vector<vector<int>>(n+1);
  60.         for(auto [a, b] : edges){
  61.             adj[a].pb(b);
  62.             adj[b].pb(a);
  63.         }
  64.         dfs(1, 0);
  65.         head[1] = 1;
  66.         dfs2(1, 0);
  67.         cerr << "pr tree\n";
  68.         for(int i = 1; i<=n; i++){
  69.             cerr << i << " " << par[i] << " " << head[i] << "\n";
  70.         }
  71.     }
  72.     void dfs(int x, int p){
  73.         sz[x] = 1;
  74.         dep[x] = dep[p] + 1;
  75.         par[x] = p;
  76.         for(auto &i : adj[x]){
  77.             if(i == p) continue;
  78.             dfs(i, x);
  79.             sz[x] += sz[i];
  80.             if(adj[x][0] == p || sz[i] > sz[adj[x][0]]) swap(adj[x][0], i);
  81.         }
  82.     }
  83.     void dfs2(int x, int p){
  84.         tin[x] = ind++;
  85.         for(auto &i : adj[x]){
  86.             if(i == p) continue;
  87.             head[i] = (i == adj[x][0] ? head[x] : i);
  88.             dfs2(i, x);
  89.         }
  90.         out[x] = ind;
  91.     }
  92.     int lca(int a, int b){
  93.         while(head[a] != head[b]){
  94.             if(dep[head[a]] > dep[head[b]]) swap(a, b);
  95.             b = par[head[b]];
  96.         }
  97.         if(dep[a] > dep[b]) swap(a, b);
  98.         return a;
  99.     }
  100. };
  101.  
  102. const int s = (1 << 18);
  103.  
  104. hld A, B;
  105. vector<pair<int, int>> edges1, edges2;
  106. vector<int> tour;
  107. int n, q, P;
  108. int root[MX];
  109.  
  110. const int NN = MX * 60;
  111.  
  112. int lc[NN], rc[NN], sum[NN];
  113. int ind = 0;
  114.  
  115. bool cmp(int a, int b){
  116.     return A.tin[a] < A.tin[b];
  117. }  
  118.  
  119. void dup(int& k){
  120.     ind++;
  121.     lc[ind] = lc[k];
  122.     rc[ind] = rc[k];
  123.     sum[ind] = sum[k];
  124.     k = ind;
  125. }
  126.  
  127. void U(int p, int& k, int L, int R){
  128.     dup(k);
  129.     if(L + 1 == R){
  130.         sum[k]++;
  131.         return ;
  132.     }
  133.     int mid = (L + R)/2;
  134.     if(p < mid) U(p, lc[k], L, mid);
  135.     else U(p, rc[k], mid, R);
  136.     sum[k] = sum[lc[k]] + sum[rc[k]];
  137. }
  138.  
  139. int S(int qL, int qR, int k, int L, int R){
  140.     if(qR <= L || R <= qL) return 0;
  141.     if(qL <= L && R <= qR){
  142.         return sum[k];
  143.     }
  144.     int mid = (L + R)/2;
  145.     return S(qL, qR, lc[k], L, mid) + S(qL, qR, rc[k], mid, R);
  146. }
  147.  
  148. int query(int a, int u, int v){
  149.     int ans = 0;
  150.     //cerr << u << " " << v << "\n";
  151.     //int cnt = 0;
  152.     while(B.head[u] != B.head[v]){
  153.     //cnt++;
  154.     //if(cnt > 10) break;
  155.         //cerr << u << " " << v << "\n";
  156.         if(B.dep[B.head[u]] < B.dep[B.head[v]]) swap(u, v);
  157.         ans += S(B.tin[B.head[u]], B.tin[u]+1, root[a], 0, s);
  158.         u = B.par[B.head[u]];
  159.     }
  160.     if(B.dep[u] > B.dep[v]) swap(u, v);
  161.     ans += S(B.tin[u], B.tin[v]+1, root[a], 0, s);
  162.     return ans;
  163. }
  164.  
  165. #define res(a) ((a + last*P - 1)%n +1)
  166.  
  167. void solve(){
  168.     n = in, q = in, P = in;
  169.     for(int i = 1; i<n; i++){
  170.         int a =in, b = in;
  171.         edges1.pb(mp(a, b));
  172.     }
  173.     for(int i = 1; i<n; i++){
  174.         int a =in, b=  in;
  175.         edges2.pb(mp(a, b));
  176.     }  
  177.     cerr << "\n";
  178.     A = hld(edges1);
  179.     B = hld(edges2);
  180.     tour = vector<int>(n);
  181.     iota(all(tour), 1);
  182.     sort(all(tour), cmp);
  183.     cerr << "\n";
  184.     for(int x : tour){
  185.         root[x] = root[A.par[x]];
  186.         cerr << x << " ";
  187.         U(B.tin[x], root[x], 0, s);
  188.     }
  189.     cerr << "\n";
  190.     int last = 0;
  191.     for(int qaq = 1; qaq <= q; qaq++){
  192.         int a = in, b = in, c = in, d = in;
  193.         a = res(a);
  194.         b = res(b);
  195.         c = res(c);
  196.         d = res(d);
  197.         cerr << a << " " << b << " " << c << " " << d << "\n";
  198.         int l = A.lca(a, b), p = A.par[l];
  199.         int a1 = query(a, c, d);
  200.         int a2 = query(b, c, d);
  201.         int a3 = query(l, c, d);
  202.         int a4 = query(p, c, d);
  203.         cerr << a << " " << a1 << "\n";
  204.         cerr << b << " " << a2 << "\n";
  205.         cerr << l << " " << a3 << "\n";
  206.         cerr << p << " " << a4 << "\n";
  207.         last = a1 + a2 - a3 - a4;
  208.         cout << last << "\n";
  209.     }
  210.    
  211.  
  212.  
  213.  
  214. }
  215.  
  216. signed main(){
  217.   cin.tie(0) -> sync_with_stdio(0);
  218.  
  219.   int T = 1;
  220.   //cin >> T;
  221.   for(int i = 1; i<=T; i++){
  222.         solve();
  223.     }
  224.   return 0;
  225. }
  226.  
  227.  
  228.  
Advertisement
Add Comment
Please, Sign In to add comment