做法是O(wnnlogn)的,代码如下。
#include<bits/stdc++.h>
#include<cmath>
#define ll long long
#define ull unsigned long long
#define ld long double
#define N 1000010
#define For(i,a,b) for(int i=a;i<=b;i++)
#define Rof(i,a,b) for(int i=a;i>=b;i--)
#define ls x<<1
#define rs x<<1|1
#define lson ls,l,mid
#define rson rs,mid+1,r
#define pb push_back
#define mk make_pair
#define pii pair<ll,ll>
#define pque priority_queue
using namespace std;
vector<int >e[N];
int f[N],sz[N];
ll n,ans=0;
ll read(){
ll x=0,f=1;char ch=getchar();
while(ch<'0'||ch>'9'){if(ch=='-')f=-1;ch=getchar();}
while(ch>='0'&&ch<='9'){x=x*10+ch-'0';ch=getchar();}
return x*f;
}
bitset<64 >b64;
bitset<128 >b128;
bitset<256 >b256;
bitset<512 >b512;
bitset<1024 >b1024;
bitset<2048 >b2048;
bitset<4096 >b4096;
bitset<8192 >b8192;
bitset<16384 >b16384;
bitset<32768 >b32768;
bitset<65536 >b65536;
bitset<131072 >b131072;
bitset<262144 >b262144;
bitset<524288 >b524288;
bitset<1048576 >b1048576;
vector<int >num;
int s[N];
void sol(int u,int fa){
if(sz[u]==1) return;
for(auto v:e[u]) s[sz[v]]++;
for(auto v:e[u]){
if(s[sz[v]]){
for(int j=1;j<=s[sz[v]];j<<=1){
s[sz[v]]-=j;
num.pb(j*sz[v]);
}
if(s[sz[v]]) num.pb(s[sz[v]]*sz[v]),s[sz[v]]=0;
}
}
if(sz[u]<=64){
b64=0;
b64[0]=1;
for(auto v:num) b64|=(b64<<v);
ll mx=0;
For(i,0,63) if(b64[i]) mx=max(mx,(ll)i*(sz[u]-i-1));
ans+=mx;
}else if(sz[u]<=128){
b128=0;
b128[0]=1;
for(auto v:num) b128|=(b128<<v);
ll mx=0;
For(i,0,127) if(b128[i]) mx=max(mx,(ll)i*(sz[u]-i-1));
ans+=mx;
}else if(sz[u]<=256){
b256=0;
b256[0]=1;
for(auto v:num) b256|=(b256<<v);
ll mx=0;
For(i,0,255) if(b256[i]) mx=max(mx,(ll)i*(sz[u]-i-1));
ans+=mx;
}else if(sz[u]<=512){
b512=0;
b512[0]=1;
for(auto v:num) b512|=(b512<<v);
ll mx=0;
For(i,0,511) if(b512[i]) mx=max(mx,(ll)i*(sz[u]-i-1));
ans+=mx;
}else if(sz[u]<=1024){
b1024=0;
b1024[0]=1;
for(auto v:num) b1024|=(b1024<<v);
ll mx=0;
For(i,0,1023) if(b1024[i]) mx=max(mx,(ll)i*(sz[u]-i-1));
ans+=mx;
}else if(sz[u]<=2048){
b2048=0;
b2048[0]=1;
for(auto v:num) b2048|=(b2048<<v);
ll mx=0;
For(i,0,2047) if(b2048[i]) mx=max(mx,(ll)i*(sz[u]-i-1));
ans+=mx;
}else if(sz[u]<=4096){
b4096=0;
b4096[0]=1;
for(auto v:num) b4096|=(b4096<<v);
ll mx=0;
For(i,0,4095) if(b4096[i]) mx=max(mx,(ll)i*(sz[u]-i-1));
ans+=mx;
}else if(sz[u]<=8192){
b8192=0;
b8192[0]=1;
for(auto v:num) b8192|=(b8192<<v);
ll mx=0;
For(i,0,8191) if(b8192[i]) mx=max(mx,(ll)i*(sz[u]-i-1));
ans+=mx;
}else if(sz[u]<=16384){
b16384=0;
b16384[0]=1;
for(auto v:num) b16384|=(b16384<<v);
ll mx=0;
For(i,0,16383) if(b16384[i]) mx=max(mx,(ll)i*(sz[u]-i-1));
ans+=mx;
}else if(sz[u]<=32768){
b32768=0;
b32768[0]=1;
for(auto v:num) b32768|=(b32768<<v);
ll mx=0;
For(i,0,32767) if(b32768[i]) mx=max(mx,(ll)i*(sz[u]-i-1));
ans+=mx;
}else if(sz[u]<=65536){
b65536=0;
b65536[0]=1;
for(auto v:num) b65536|=(b65536<<v);
ll mx=0;
For(i,0,65535) if(b65536[i]) mx=max(mx,(ll)i*(sz[u]-i-1));
ans+=mx;
}else if(sz[u]<=131072){
b131072=0;
b131072[0]=1;
for(auto v:num) b131072|=(b131072<<v);
ll mx=0;
For(i,0,131071) if(b131072[i]) mx=max(mx,(ll)i*(sz[u]-i-1));
ans+=mx;
}else if(sz[u]<=262144){
b262144=0;
b262144[0]=1;
for(auto v:num) b262144|=(b262144<<v);
ll mx=0;
For(i,0,262143) if(b262144[i]) mx=max(mx,(ll)i*(sz[u]-i-1));
ans+=mx;
}else if(sz[u]<=524288){
b524288=0;
b524288[0]=1;
for(auto v:num) b524288|=(b524288<<v);
ll mx=0;
For(i,0,524287) if(b524288[i]) mx=max(mx,(ll)i*(sz[u]-i-1));
ans+=mx;
}else{
b1048576=0;
b1048576[0]=1;
for(auto v:num) b1048576|=(b1048576<<v);
ll mx=0;
For(i,0,1048575) if(b1048576[i]) mx=max(mx,(ll)i*(sz[u]-i-1));
ans+=mx;
}
while(!num.empty()) num.pop_back();
for(auto v:e[u]) sol(v,u);
}
int main()
{
n=read();
For(i,1,n) sz[i]=1;
For(i,2,n){
f[i]=read();
e[f[i]].pb(i);
}
Rof(i,n,2) sz[f[i]]+=sz[i];
sol(1,0);
cout<<ans;
return 0;
}
遇到以下数据会RE。
8
1 2 3 4 5 6 7