以下是代码部分
以下是代码部分
#include<iostream>
using namespace std;
const int maxn = 5005;
int head[maxn];
int tail[maxn * 2];
int next[maxn * 2];
int son[maxn];
int m = 1;
int d[maxn];
int F[maxn];
int G[maxn];
int* f[maxn];
int* g[maxn];
void link(int a, int b) {
m++;
tail[m] = b;
next[m] = head[a];
head[a] = m;
}
void dfs1(int u, int fa) {
d[u] = 1;
for (int e = head[u]; e != 0; e = next[e]) {
int v = tail[e];
if (v != fa) {
dfs1(v, u);
if (d[u] < d[v] + 1) {
son[u] = v;
d[u] = d[v] + 1;
}
}
}
}
long long ans = 0;
void dfs(int u, int fa, int pt) {
f[u] = F + pt;
g[u] = G + pt;
int rr = pt;
pt += d[son[u]] + 1;
f[u][0] = 1;
for (int e = head[u]; e != 0; e = next[e]) {
int v = tail[e];
if (v == son[u]) continue;
if (v != fa) {
dfs(v, u, pt);
pt += d[v];
}
}
if (son[u] > 0)
dfs(son[u], u, rr + 1);
for (int i = 0; i + 1 < d[son[u]]; ++i) {
g[u][i] = g[son[u]][i + 1];
}
for (int e = head[u]; e != 0; e = next[e]) {
int v = tail[e];
if (v != fa && v != son[u]) {
for (int i = 0; i < d[v]; ++i) {
ans += (long long)g[u][i+1] * f[v][i];
}
for (int i = 0; i < d[v]; ++i){
ans += (long long)f[u][i] * g[v][i+1];
}
for (int i = 0; i < d[v]; ++i) {
g[u][i+1] += f[v][i] * f[u][i+1];
}
for (int i = 0; i + 1 < d[v]; ++i) {
g[u][i] += g[v][i+1];
}
for (int i = 0; i < d[v]; ++i) {
f[u][i+1] += f[v][i];
}
}
}
for (int i = rr + d[u]; i < pt; i++) {
f[u][i] = g[u][i] = 0;
}
}
int main() {
int n;
cin >> n;
for (int i = 1; i < n; ++i) {
int a, b;
cin >> a >> b;
link(a, b);
link(b, a);
}
dfs1(1, 0);
dfs(1, 0, 1);
cout << ans;
}
样例过不了