有无大佬帮我们看看T3代码:(这三个代码乍一看都是O(n)的,但正确性未知……)
#include <cstdio>
#include <iostream>
#include <vector>
using namespace std;
int n, x;
int minn;
int val[100005],fmx[100005],fm[100005], js[100005];
struct mode
{
int sum_of_son;
int maxson;
vector<int> son;
} s[100005];
inline int max(int x, int y)
{
if(x>y)
return x;
return y;
}
inline void dpfs(int root)
{
js[root]++;
if(s[root].sum_of_son==0)
{
fmx[root]=val[root];
fm[root]=1;
return;
}
for(int i=0;i<s[root].sum_of_son;++i)
{
dpfs(s[root].son[i]);
fm[root]+=fm[s[root].son[i]];
fmx[root]=max(fmx[root],fmx[s[root].son[i]]);
}
if(fm[root]*fmx[root]>val[root])
fm[root]=1,fmx[root]=val[root];
return;
}
int main()
{
// freopen("in.txt","r",stdin);
// freopen("line.out","w",stdout);
scanf("%d",&n);
for (int i = 1; i <= n; ++i)
scanf("%d",&val[i]);
for (int i = 2; i <= n; ++i)
{
scanf("%d",&x);
++s[x].sum_of_son;
s[x].son.push_back(i);
if (x > s[x].maxson)
s[x].maxson = x;
}
dpfs(1);
printf("%d", fm[1]*fmx[1]);
for(int i = 1; i <= n; i++) {
printf("%d ", js[i]);
}
return 0;
}
#include<bits/stdc++.h>
#define LL long long
using namespace std;
const int N=1e5+5;
int n,head[N],nxt[N],To[N],cntm,w[N],fax[N],fs[N];
LL f[N];
void add(int u,int v)
{
++cntm;
nxt[cntm]=head[u];
head[u]=cntm;
To[cntm]=v;
}
void dfs(int x)
{
int s=0,maxn=0;
LL xc;
f[x]=w[x];
for(int i=head[x];i;i=nxt[i])
{
dfs(To[i]);
if(fs[To[i]]==1)
maxn=max(maxn,w[To[i]]);
else
maxn=max(maxn,fax[To[i]]);
s+=fs[To[i]];
}
xc=s*maxn;
if(xc<f[x]&&s!=0)
{
f[x]=xc;
fs[x]=s;
}
else
fs[x]=1;
fax[x]=maxn;
return ;
}
int main()
{
// freopen("line.in","r",stdin);
// freopen("line.out","w",stdout);
scanf("%d",&n);
for(int i=1;i<=n;++i)
scanf("%lld",&w[i]);
for(int i=2,fa;i<=n;++i)
{
scanf("%d",&fa);
add(fa,i);
}
dfs(1);
printf("%lld",f[1]);
}
#include <stdio.h>
const int MAXN = 1e5 + 10;
const long long INF = 1e15 * 2 + 150;
long long w[MAXN], tem[MAXN], hb[MAXN];
int head[MAXN], nx[MAXN], to[MAXN], tot;
long long anss = INF;
void Add(int x, int y) {
tot++;
nx[tot] = head[x];
to[tot] = y;
head[x] = tot;
}
long long Dfs(int x, long long upw) {
if(w[x] <= upw) return 1;
if(!head[x]) return -1;
long long ans = 0;
for(int i = head[x]; i; i = nx[i]) {
long long ret = Dfs(to[i], upw);
if(ret > 0) {
ans += ret;
if(ans * upw >= anss) return -1;
}
else {
return -1;
}
}
if(ans * upw >= anss) return -1;
return ans;
}
long long min(long long a, long long b) {
return a < b ? a : b;
}
void Msort(int l, int r) {
if(l == r) return;
int mid = (l + r) / 2;
Msort(l, mid);
Msort(mid + 1, r);
int i = l, j = mid + 1, k = l;
while(i <= mid && j <= r) {
if(hb[i] < hb[j]) {
tem[k] = hb[i];
i++;
k++;
}
else {
tem[k] = hb[j];
j++;
k++;
}
}
while(i <= mid) {
tem[k] = hb[i];
i++;
k++;
}
while(j <= r) {
tem[k] = hb[j];
j++;
k++;
}
for(int x = l; x <= r; x++) {
hb[x] = tem[x];
}
}
int main() {
// freopen("in.txt", "r", stdin);
// freopen("line.out", "w", stdout);
int n;
scanf("%d", &n);
for(int i = 1; i <= n; i++) {
scanf("%lld", &w[i]);
hb[i] = w[i];
}
for(int i = 2; i <= n; i++) {
int t;
scanf("%d", &t);
Add(t, i);
}
Msort(1, n);
long long last = -1;
for(int i = 1; i <= n; i++) {
if(hb[i] == last || hb[i] >= anss) continue;
long long ret_num = Dfs(1, hb[i]);
// printf("%lld: %lld\n", hb[i], ret_num);
if(ret_num > 0) {
anss = min(anss, ret_num * hb[i]);
}
last = hb[i];
}
printf("%lld", anss);
return 0;
}