Not a member of Pastebin yet?
Sign Up,
it unlocks many cool features!
- //#pragma optimization_level 3
- //#pragma GCC target("sse,sse2,sse3,ssse3,sse4,popcnt,abm,mmx,avx,avx2,tune=native")
- //#pragma GCC optimize("Ofast,no-stack-protector,unroll-loops,fast-math")
- #include <iostream>
- #include <algorithm>
- #include <fstream>
- #include <vector>
- #include <queue>
- #include <functional>
- #include <set>
- #include <map>
- #include <math.h>
- #include <cmath>
- #include <string>
- #include <random>
- #include <unordered_set>
- #include <unordered_map>
- #include <bitset>
- #include <string.h>
- #include <stack>
- #include <assert.h>
- #include <list>
- #include <time.h>
- #include <memory>
- #include <chrono>
- using namespace std;
- //
- #define fast cin.tie(0);cout.tie(0);cin.sync_with_stdio(0);cout.sync_with_stdio(0);
- #define cin in
- //#define cout out
- #define ll long long
- #define db double
- #define ld long double
- #define uset unordered_set
- #define umap unordered_map
- #define ms multiset
- #define pb push_back
- #define pq priority_queue
- #define umap unordered_map
- #define uset unordered_set
- #define ull unsigned long long
- #define pii pair<int, int>
- #define pll pair<ll, ll>
- #define pdd pair<ld, ld>
- #define pnn pair<Node*, Node*>
- #define uid uniform_int_distribution
- #define PI acos(-1.0)
- //#define sort(a, b) sort(a.begin(), a.end(), b())
- //mt19937 rnd(chrono::steady_clock::now().time_since_epoch().count());
- ifstream in("input.txt");
- ofstream out("output.txt");
- const int MAX_N = 1e5;
- struct CentInfo {
- int cent, d, son;
- };
- //for vertices
- vector<CentInfo> cents[MAX_N];
- //for centroids
- vector<umap<int, int>> c_dists[MAX_N];
- umap<int, int> c_sumDists[MAX_N];
- int sub_sz[MAX_N];
- bool used[MAX_N];
- int up_d_dist[MAX_N];
- vector<int> g[MAX_N];
- int n, D;
- //calculates values, v - son of cent initially
- void calc_vals(int v, int pr, int cent, int d, int son_num) {
- cents[v].push_back({cent, d, son_num});
- c_dists[cent][son_num][d]++;
- c_sumDists[cent][d]++;
- for (int to : g[v]) {
- if (used[to] || to == pr) continue;
- calc_vals(to, v, cent, d + 1, son_num);
- }
- }
- //returns size of v's component and calculates sizes of subtrees
- int get_sz(int v, int pr) {
- sub_sz[v] = 1;
- for (int to : g[v]) {
- if (used[to] || to == pr) continue;
- sub_sz[v] += get_sz(to, v);
- }
- return sub_sz[v];
- }
- //d = 0 case ?
- //calculates values for centroid = "cent"
- void calc_vals(int cent) {
- c_dists[cent].resize(g[cent].size());
- get_sz(cent, -1);
- for (int i = 0; i < (int)g[cent].size(); i++)
- if(!used[g[cent][i]])
- calc_vals(g[cent][i], cent, cent, 1, i);
- }
- int get_cent(int v, int sz) {
- for (int to : g[v]) {
- if (used[to]) continue;
- if (2 * sub_sz[to] >= sz)
- return get_cent(to, sz);
- }
- return v;
- }
- void build_cents(int v) {
- v = get_cent(v, get_sz(v, -1));
- calc_vals(v);
- used[v] = 1;
- for (int to : g[v]) {
- if (!used[to])
- build_cents(to);
- }
- }
- //returns total number of d-distanced vertices from 'v'
- int get_d_dist(int v) {
- int ans = (c_sumDists[v].count(D) ? c_sumDists[v][D] : 0);
- for (auto [cent, d, son] : cents[v]) {
- int left = D - d;
- if (left < 0) continue;
- else if (left == 0) ++ans;
- else {
- int tmp = (c_sumDists[cent].count(left) ? c_sumDists[cent][left] : 0) -
- (c_dists[cent][son].count(left) ? c_dists[cent][son][left] : 0);
- ans += tmp;
- }
- }
- return ans;
- }
- void input() {
- cin >> n >> D;
- for (int i = 0; i < n - 1; ++i) {
- int a, b;
- cin >> a >> b;
- --a; --b;
- g[a].push_back(b);
- g[b].push_back(a);
- }
- }
- int main() {
- input();
- if (D % 2) {
- cout << 0;
- return 0;
- }
- D /= 2;
- build_cents(0);
- for (int v = 0; v < n; v++)
- up_d_dist[v] = get_d_dist(v);
- }
Advertisement
Add Comment
Please, Sign In to add comment