和题解一个做法。
#include <bits/stdc++.h>
using namespace std;
#define FOR(i,a,b) for(int i=a;i<=b;i++)
#define REP(i,a,b) for(int i=a;i>=b;i--)
#define pb push_back()
#define mkpr make_pair
typedef long long ll;
typedef pair<int,int> pii;
typedef pair<ll,ll> pll;
const int maxn=1e6+5;
int head[maxn];
struct EDGE
{
int to,nxt;
}edge[maxn<<1];
int cnt;
void add(int u,int to)
{
edge[++cnt].to=to;
edge[cnt].nxt=head[u];
head[u]=cnt;
}
int n;
int p[maxn];
ll dp[maxn];
int siz[maxn];
template<int len=1>
void getbs(int ned,vector<int>& a,int u)
{
if(len<ned)
{
getbs<min(len*2,maxn)>(ned,a,u);
return;
}
bitset<len> dp2;
dp2.reset();
dp2[0]=1;
for(int i=0;i<a.size();i++)
{
// printf("a[%d] = %d\n",i,a[i]);
dp2|=(dp2<<a[i]);
}
int minn=1e9;
int aa=0;
for(int i=0;i<siz[u];i++)
{
if(dp2[i])
{
// printf("can: %d\n",i);
if(abs(siz[u]-1-i-i)<=minn)
{
minn=abs(siz[u]-1-i-i);
aa=i;
}
}
}
dp[u]+=ll(aa)*ll(siz[u]-1-aa);
// printf("sec dp[%d] = %lld aa = %d siz = %d\n",u,dp[u],aa,siz[u]);
}
int cntsize[maxn];
void dfs(int u,int _fa)
{
siz[u]=1;
vector<int> a;
for(int i=head[u];i;i=edge[i].nxt)
{
int to=edge[i].to;
if(to==_fa)continue;
dfs(to,u);
dp[u]+=dp[to];
siz[u]+=siz[to];
cntsize[siz[to]]++;
}
int stop=0;
for(int i=head[u];i;i=edge[i].nxt)
{
int to=edge[i].to;
if(to==_fa)continue;
if(siz[to]>=siz[u]-1-siz[to])
{
stop=siz[to];
break;
}
}
if(stop)
{
dp[u]+=ll(stop)*ll(siz[u]-1-stop);
for(int i=head[u];i;i=edge[i].nxt)
{
int to=edge[i].to;
if(to==_fa)continue;
cntsize[siz[to]]=0;
}
// printf("pre dp[%d] = %lld\n",u,dp[u]);
return;
}
for(int i=head[u];i;i=edge[i].nxt)
{
int to=edge[i].to;
if(to==_fa)continue;
// if(!cntsize[siz[to]])continue;
int base=1;
while(base<cntsize[siz[to]])
{
cntsize[siz[to]]-=base;
a.emplace_back((ll)base*(ll)siz[to]);
base*=2;
}
a.emplace_back((ll)cntsize[siz[to]]*(ll)siz[to]);
cntsize[siz[to]]=0;
// a.emplace_back(siz[to]);
}
getbs<1>(siz[u],a,u);
}
signed main()
{
scanf("%d",&n);
FOR(i,2,n)
{
scanf("%d",&p[i]);
add(i,p[i]);
add(p[i],i);
}
dfs(1,0);
printf("%lld",dp[1]);
return 0;
}