警示后人(如果你95分TLE#9)
查看原帖
警示后人(如果你95分TLE#9)
565040
Conan15楼主2023/6/24 14:42

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;
}
2023/6/24 14:42
加载中...