Not a member of Pastebin yet?
Sign Up,
it unlocks many cool features!
- #include <bits/stdc++.h>
- #define ll long long
- using namespace std;
- const int N = 2e5 + 5;
- int n, m;
- vector<int> a[N];
- int d[N], dem[N];
- bool ok = false;
- ll ans;
- ll rood(int u) {
- if(dem[u] > 1) return u;
- if(dem[u] == 1) {
- int uv = a[u][0];
- rood(uv);
- }
- }
- void DFS(int u) {
- for(int v : a[u]) {
- if(v != 0) {
- d[v] = d[u] + 1;
- DFS(v);
- }
- }
- }
- void solve() {
- for(int i = 2; i <= n; ++i) {
- int x; cin >> x;
- a[x].push_back(i);
- dem[x]++;
- if(dem[x] >= 2) ok = true;
- }
- if(!ok) {
- cout << n * m << '\n';
- return;
- }
- if(dem[1] == n-1) {
- cout << 3*m << '\n';
- return;
- }
- int maxx = 0, minn = 0;
- ll uv = 1;
- uv = max(uv, rood(1));
- for(int i = 1; i <= n; ++i)
- a[i].push_back(0);
- for(int u : a[uv]) {
- if(a[u].size() != 1) {
- fill(d + 1, d + n + 1, 0);
- DFS(u);
- int k = *max_element(d + 1, d + n + 1);
- if(maxx < k) {
- minn = maxx;
- maxx = k;
- }
- else if(minn < k) {
- minn = k;
- }
- }
- }
- maxx += 2; minn++;
- //cout << uv << ' ' << maxx << ' ' << minn << '\n';
- ans = (maxx*m + (uv-1) * (m-1)) + (minn + (maxx+uv-1) * (m-1));
- cout << ans << '\n';
- }
- int main()
- {
- freopen("in.txt", "r", stdin);
- //freopen("FIREWORKS.inp", "r", stdin);
- //freopen("FIREWORKS.out", "w", stdout);
- ios_base::sync_with_stdio(false);
- cin.tie(0); cout.tie(0);
- cin >> n >> m;
- solve();
- return 0;
- }
Advertisement
Add Comment
Please, Sign In to add comment