95分代码:
#include <bits/stdc++.h>
using namespace std;
const long long N = 2e6 + 10, M = N * 2;
long long h[N], e[M], w[M], ne[M], idx;
void add(long long a, long long b, long long c) {
e[++idx] = b, w[idx] = c, ne[idx] = h[a], h[a] = idx;
}
bool st[N];
long long top = 0;
pair<long long, long long> stk[N];
long long flag = 0;
//?
long long id[N], cnt = 0; //?????
bool cir[N]; //???
long long p[N];
//?
void find(long long u, long long edge) {
st[u] = true; stk[++top] = make_pair(u, w[edge]);
for (long long i = h[u]; i; i = ne[i]) {
long long v = e[i];
if ((i ^ 1) == edge || flag) continue;
if (st[v]) {
while (stk[top].first != v) {
long long s = stk[top].first;
id[++cnt] = s, cir[s] = true, p[cnt] = stk[top].second, top--;
}
id[++cnt] = v, cir[v] = true, p[cnt] = w[i];
flag = true; //???
return;
}
find(v, i);
}
top--;
}
long long dis = 0, dis_u = 0;
bool QwQ[N];
void dfs(long long u, long long edge, long long D) {
if (dis < D) dis = D, dis_u = u;
for (long long i = h[u]; i; i = ne[i]) {
long long v = e[i];
if ((i ^ 1) == edge || cir[v]) continue;
dfs(v, i, w[i] + D);
}
}
long long dist(long long u) {
dis = dis_u = 0;
dfs(u, 0, 0);
dis = 0;
dfs(dis_u, 0, 0);
return dis;
}
long long d[N], dd[N];
long long n;
long long Max_1[N], Max_2[N];
long long front[N], back[N];
long long solve(long long start) {
dis = dis_u = 0;
memset(Max_1, 0, sizeof Max_1);
memset(Max_2, 0, sizeof Max_2);
memset(front, 0, sizeof front);
memset(back, 0, sizeof back);
memset(d, 0, sizeof d);
memset(dd, 0, sizeof dd);
memset(p, 0, sizeof p);
memset(id, 0, sizeof id);
memset(cir, 0, sizeof cir);
memset(st, 0, sizeof st);
flag = cnt = top = 0;
find(start, 0);
// cout << "!!!!!!!!!!!!!!!!" << cnt << endl;
for (long long i = 1; i <= cnt; i++) {
cir[id[i]] = false; d[i] = dist(id[i]); cir[id[i]] = true;
dis = 0; dfs(id[i], 0, 0); dd[i] = dis;
}
// cout << '\t'; for (long long i = 1; i <= cnt; i++) cout << d[i] << ' '; puts("");
long long Dis = 0, Len = 0;
for (long long i = 1; i <= cnt; i++) {
front[i] = max(front[i - 1], dd[i] + Dis);
if (i > 1) Max_1[i] = max(Max_1[i - 1], Len + Dis + dd[i]);
Len = max(Len, dd[i] - Dis), Dis += p[i];
}
long long ans = 0;
dis = Len = 0, ans = Max_1[cnt];
for (long long i = cnt; i >= 1; i--) {
back[i] = max(back[i + 1], dd[i] + dis);
if (i < cnt) Max_2[i] = max(Max_2[i + 1], Len + dis + dd[i]);
Len = max(Len, dd[i] - dis), dis += p[i - 1];
}
for (long long i = cnt - 1; i >= 1; i--)
ans = max(ans, max(max(Max_1[i], Max_2[i + 1]), front[i] + back[i + 1] + p[cnt]));
for (long long i = 1; i <= cnt; i++) ans = max(ans, d[i]);
return ans;
}
void dfsqwq(long long u) {
QwQ[u] = true;
for (long long i = h[u]; i; i = ne[i]) {
long long v = e[i];
if (QwQ[v]) continue;
dfsqwq(v);
}
}
int main() {
idx++;
scanf("%lld", &n);
for (long long i = 1; i <= n; i++) {
long long a, b; scanf("%lld%lld", &a, &b);
add(i, a, b); add(a, i, b);
}
long long awa = 0;
for (long long i = 1; i <= n; i++) {
if (QwQ[i]) continue;
dfsqwq(i);
// for (long long j = 1; j <= n; j++) cout << QwQ[j] << ' '; puts("");
long long res = solve(i);
// cout << i << ' ' << res << endl;
awa += 1ll * res;
// if (!QwQ[i]) awa += solve(i);
}
printf("%lld\n", awa);
return 0;
}
由于要对每个连通块,即每个基环树进行计算直径,需要在 solve 函数里进行数组和全局变量的初始化。
但我在上面代码中用了memset……
“memset 的时间复杂度为 $O(n)。对于Educational Codeforces Round 84 (Rated for Div. 2) 的 b 题,经测试定义一个数组使用memset的时间为 1200ms,而定义两个数组使用 memset 的时间就提示超时。最后换成循环初始化的时候 AC 时间仅为 93ms。所以对于测试数据很多的情况还是慎用memset。”
这段话来自这里。
因此还是要用 for 循环。
修改后 AC 代码如下:
加入 register 和快读是为了优化时间复杂度。
#include <bits/stdc++.h>
using namespace std;
int read() {
int s = 0, f = 1;
char ch = getchar();
while(!isdigit(ch)) {
if(ch == '-') f = -1;
ch = getchar();
}
while(isdigit(ch)) {
s = s * 10 + ch - '0';
ch = getchar();
}
return s * f;
}
const long long N = 2e6 + 10, M = N * 2;
int h[N], e[M], w[M], ne[M], idx;
void add(int a, int b, int c) {
e[++idx] = b, w[idx] = c, ne[idx] = h[a], h[a] = idx;
}
bool st[N];
int top = 0;
pair<long long, long long> stk[N];
bool flag = 0;
//?
int id[N], cnt = 0; //?????
bool cir[N]; //???
long long p[N];
//?
void find(long long u, long long edge) {
st[u] = true; stk[++top] = make_pair(u, w[edge]);
for (register int i = h[u]; i; i = ne[i]) {
long long v = e[i];
if ((i ^ 1) == edge || flag) continue;
if (st[v]) {
while (stk[top].first != v) {
long long s = stk[top].first;
id[++cnt] = s, cir[s] = true, p[cnt] = stk[top].second, top--;
}
id[++cnt] = v, cir[v] = true, p[cnt] = w[i];
flag = true; //???
return;
}
find(v, i);
}
top--;
}
long long dis = 0, dis_u = 0;
bool QwQ[N];
void dfs(long long u, long long edge, long long D) {
if (dis < D) dis = D, dis_u = u;
for (register int i = h[u]; i; i = ne[i]) {
long long v = e[i];
if ((i ^ 1) == edge || cir[v]) continue;
dfs(v, i, w[i] + D);
}
}
long long dist(long long u) {
dis = dis_u = 0;
dfs(u, 0, 0);
dis = 0;
dfs(dis_u, 0, 0);
return dis;
}
long long d[N], dd[N];
long long n;
long long Max_1[N], Max_2[N];
long long front[N], back[N];
long long solve(long long start) {
find(start, 0);
// cout << "!!!!!!!!!!!!!!!!" << cnt << endl;
for (register int i = 1; i <= cnt; i++) {
cir[id[i]] = false; d[i] = dist(id[i]); cir[id[i]] = true;
dis = 0; dfs(id[i], 0, 0); dd[i] = dis;
}
// cout << '\t'; for (register int i = 1; i <= cnt; i++) cout << d[i] << ' '; puts("");
long long Dis = 0, Len = 0;
for (register int i = 1; i <= cnt; i++) {
front[i] = max(front[i - 1], dd[i] + Dis);
if (i > 1) Max_1[i] = max(Max_1[i - 1], Len + Dis + dd[i]);
Len = max(Len, dd[i] - Dis), Dis += p[i];
}
long long ans = 0;
dis = Len = 0, ans = Max_1[cnt];
for (register int i = cnt; i >= 1; i--) {
back[i] = max(back[i + 1], dd[i] + dis);
if (i < cnt) Max_2[i] = max(Max_2[i + 1], Len + dis + dd[i]);
Len = max(Len, dd[i] - dis), dis += p[i - 1];
}
for (register int i = cnt - 1; i >= 1; i--)
ans = max(ans, max(max(Max_1[i], Max_2[i + 1]), front[i] + back[i + 1] + p[cnt]));
for (register int i = 1; i <= cnt; i++) ans = max(ans, d[i]);
for (int i = 1; i <= cnt; i++) cir[id[i]] = id[i] = 0;
for (int i = 0; i <= cnt + 5; i++)
Max_1[i] = Max_2[i] = front[i] = back[i] = d[i] = dd[i] = p[i] = 0;
for (int i = 1; i <= n; i++) st[i] = 0;
dis = dis_u = 0;
flag = cnt = top = 0;
return ans;
}
void dfsqwq(int u) {
QwQ[u] = true;
for (int i = h[u]; i; i = ne[i]) {
int v = e[i];
if (QwQ[v]) continue;
dfsqwq(v);
}
}
int main() {
idx++; n = read();
for (register int i = 1; i <= n; i++) {
int a, b; a = read(), b = read();
add(i, a, b); add(a, i, b);
}
long long awa = 0;
for (register int i = 1; i <= n; i++) {
if (QwQ[i]) continue;
dfsqwq(i);
// for (register int j = 1; j <= n; j++) cout << QwQ[j] << ' '; puts("");
long long res = solve(i);
// cout << i << ' ' << res << endl;
awa += 1ll * res;
// if (!QwQ[i]) awa += solve(i);
}
printf("%lld\n", awa);
return 0;
}