TrickmanOff

X-disted vertices for each in tree O(NlogN)

Aug 13th, 2020
1,351
0
Never
Not a member of Pastebin yet? Sign Up, it unlocks many cool features!
C++ 3.60 KB | None | 0 0
  1. //#pragma optimization_level 3
  2. //#pragma GCC target("sse,sse2,sse3,ssse3,sse4,popcnt,abm,mmx,avx,avx2,tune=native")
  3. //#pragma GCC optimize("Ofast,no-stack-protector,unroll-loops,fast-math")
  4. #include <iostream>
  5. #include <algorithm>
  6. #include <fstream>
  7. #include <vector>
  8. #include <queue>
  9. #include <functional>
  10. #include <set>
  11. #include <map>
  12. #include <math.h>
  13. #include <cmath>
  14. #include <string>
  15. #include <random>
  16. #include <unordered_set>
  17. #include <unordered_map>
  18. #include <bitset>
  19. #include <string.h>
  20. #include <stack>
  21. #include <assert.h>
  22. #include <list>
  23. #include <time.h>
  24. #include <memory>
  25. #include <chrono>
  26. using namespace std;
  27. //
  28. #define fast cin.tie(0);cout.tie(0);cin.sync_with_stdio(0);cout.sync_with_stdio(0);
  29. #define cin in
  30. //#define cout out
  31. #define ll long long
  32. #define db double
  33. #define ld long double
  34. #define uset unordered_set
  35. #define umap unordered_map
  36. #define ms multiset
  37. #define pb push_back
  38. #define pq priority_queue
  39. #define umap unordered_map
  40. #define uset unordered_set
  41. #define ull unsigned long long
  42. #define pii pair<int, int>
  43. #define pll pair<ll, ll>
  44. #define pdd pair<ld, ld>
  45. #define pnn pair<Node*, Node*>
  46. #define uid uniform_int_distribution
  47. #define PI acos(-1.0)
  48. //#define sort(a, b) sort(a.begin(), a.end(), b())
  49. //mt19937 rnd(chrono::steady_clock::now().time_since_epoch().count());
  50. ifstream in("input.txt");
  51. ofstream out("output.txt");
  52.  
  53. const int MAX_N = 1e5;
  54.  
  55. struct CentInfo {
  56.     int cent, d, son;
  57. };
  58.  
  59. //for vertices
  60. vector<CentInfo> cents[MAX_N];
  61. //for centroids
  62. vector<umap<int, int>> c_dists[MAX_N];
  63. umap<int, int> c_sumDists[MAX_N];
  64.  
  65. int sub_sz[MAX_N];
  66. bool used[MAX_N];
  67.  
  68. int up_d_dist[MAX_N];
  69.  
  70. vector<int> g[MAX_N];
  71. int n, D;
  72.  
  73. //calculates values, v - son of cent initially
  74. void calc_vals(int v, int pr, int cent, int d, int son_num) {
  75.     cents[v].push_back({cent, d, son_num});
  76.     c_dists[cent][son_num][d]++;
  77.     c_sumDists[cent][d]++;
  78.  
  79.     for (int to : g[v]) {
  80.         if (used[to] || to == pr) continue;
  81.         calc_vals(to, v, cent, d + 1, son_num);
  82.     }
  83. }
  84.  
  85. //returns size of v's component and calculates sizes of subtrees
  86. int get_sz(int v, int pr) {
  87.     sub_sz[v] = 1;
  88.  
  89.     for (int to : g[v]) {
  90.         if (used[to] || to == pr) continue;
  91.  
  92.         sub_sz[v] += get_sz(to, v);
  93.     }
  94.     return sub_sz[v];
  95. }
  96.  
  97. //d = 0 case ?
  98.  
  99. //calculates values for centroid = "cent"
  100. void calc_vals(int cent) {
  101.     c_dists[cent].resize(g[cent].size());
  102.     get_sz(cent, -1);
  103.     for (int i = 0; i < (int)g[cent].size(); i++)
  104.         if(!used[g[cent][i]])
  105.             calc_vals(g[cent][i], cent, cent, 1, i);
  106. }
  107.  
  108. int get_cent(int v, int sz) {
  109.     for (int to : g[v]) {
  110.         if (used[to]) continue;
  111.         if (2 * sub_sz[to] >= sz)
  112.             return get_cent(to, sz);
  113.     }
  114.     return v;
  115. }
  116.  
  117. void build_cents(int v) {
  118.     v = get_cent(v, get_sz(v, -1));
  119.     calc_vals(v);
  120.  
  121.     used[v] = 1;
  122.     for (int to : g[v]) {
  123.         if (!used[to])
  124.             build_cents(to);
  125.     }
  126. }
  127.  
  128. //returns total number of d-distanced vertices from 'v'
  129. int get_d_dist(int v) {
  130.     int ans = (c_sumDists[v].count(D) ? c_sumDists[v][D] : 0);
  131.     for (auto [cent, d, son] : cents[v]) {
  132.         int left = D - d;
  133.         if (left < 0) continue;
  134.         else if (left == 0) ++ans;
  135.         else {
  136.             int tmp = (c_sumDists[cent].count(left) ? c_sumDists[cent][left] : 0) -
  137.                 (c_dists[cent][son].count(left) ? c_dists[cent][son][left] : 0);
  138.             ans += tmp;
  139.         }
  140.     }
  141.     return ans;
  142. }
  143.  
  144. void input() {
  145.     cin >> n >> D;
  146.    
  147.     for (int i = 0; i < n - 1; ++i) {
  148.         int a, b;
  149.         cin >> a >> b;
  150.         --a; --b;
  151.         g[a].push_back(b);
  152.         g[b].push_back(a);
  153.     }
  154. }
  155.  
  156. int main() {
  157.     input();
  158.  
  159.     if (D % 2) {
  160.         cout << 0;
  161.         return 0;
  162.     }
  163.     D /= 2;
  164.  
  165.     build_cents(0);
  166.     for (int v = 0; v < n; v++)
  167.         up_d_dist[v] = get_d_dist(v);
  168. }
Advertisement
Add Comment
Please, Sign In to add comment