求助玄学问题
查看原帖
求助玄学问题
569516
C6H6楼主2023/7/24 20:06

这份代码开了 O2 之后 T 了第 10 个点, 但不开 O2 就过了,请各位大佬帮忙找错

#include <bits/stdc++.h>
using namespace std;
#define gc() (p1 == p2 && (p2 = (p1 = buf) + fread(buf, 1, MAXSIZE, stdin), p1 == p2) ? EOF : *p1++)
const int MAXSIZE = 1 << 20;
char buf[MAXSIZE], *p1, *p2; 
void read() {}
template <class T1, class ...T2>
void read(T1& ret,T2&... rest) {
    ret = 0; char c; bool f = false;
    while (!isdigit(c = gc() ) ) f = c == '-';
    while(isdigit(c) ) {
        ret = (ret << 3) + (ret << 1) + (c ^ '0');
        c = gc();
    }
    if(f) ret = -ret;
    read(rest...);
}
const int N = 250010;
struct edge {
    int to, nxt, w;
}e[N << 1];
int head[N];
int cnt = 1;
inline void add(int x, int y, int z) {
    e[++cnt] = {y, head[x], z};
    head[x] = cnt;
}
vector<int> v[N];
long long dep[N], fa[N][23], mi[N][23], dfn[N];
int a[N], tmp[N];
int k, tot;
bool vis[N];
void init(int x) {
    dfn[x] = ++tot;
    for(int i = 1; i <= 22; i++) {
        fa[x][i] = fa[fa[x][i - 1]][i - 1];
        mi[x][i] = min(mi[x][i - 1], mi[fa[x][i - 1]][i - 1]);
    }
    
    for(int i = head[x]; i; i = e[i].nxt) {
        int y = e[i].to, z = e[i].w;
        if(y == fa[x][0]) continue;
        dep[y] = dep[x] + 1;
        fa[y][0] = x;
        mi[y][0] = z;
        init(y);
    }
}
int get_lca(int x, int y) {
    if(dep[x] < dep[y]) swap(x, y);
    for(int i = 22; i >= 0; i--) if(dep[fa[x][i]] >= dep[y]) x = fa[x][i];
    if(x == y) return x;
    for(int i = 22; i >= 0; i--) if(fa[x][i] != fa[y][i]) x = fa[x][i], y = fa[y][i];
    return fa[x][0];
}
long long get_len(int x, int y) {
    int lca = get_lca(x, y);
    long long ret = 0x3f3f3f3f3f3f3f3fll;
    for(int i = 22; i >= 0; i--) if(dep[fa[x][i]] >= dep[lca]) ret = min(ret, mi[x][i]), x = fa[x][i];
    for(int i = 22; i >= 0; i--) if(dep[fa[y][i]] >= dep[lca]) ret = min(ret, mi[y][i]), y = fa[y][i];
    return ret;
}
inline bool cmp(int x, int y) {
    return dfn[x] < dfn[y];    
}
long long dfs(int x) {
    long long ret = 0;
    for(auto y : v[x]) ret += dfs(y);
    if(vis[x]) ret = get_len(1, x);
    else ret = min(get_len(1, x), ret);
    vis[x] = 0;
    v[x].clear();
    return ret;
}
int main() {
    int n;
    read(n);
    for(int i = 1, x, y, z; i < n; i++) read(x, y, z), add(x, y, z), add(y, x, z);
    memset(mi, 0x3f, sizeof(mi));
    dep[1] = 1;
    init(1);
    int m;
    read(m);
    while(m--) {
        read(k);
        for(int i = 1; i <= k; i++) read(a[i]), vis[a[i]] = 1;
        sort(a + 1, a + 1 + k, cmp);
        a[0] = 1;
        for(int i = 0, tmp = k; i <= tmp - 1; i++) a[++k] = get_lca(a[i], a[i + 1]);
        sort(a + 1, a + 1 + k, cmp);
        k = unique(a + 1, a + 1 + k) - a - 1;
        
        for(int i = 1; i < k; i++) {
            int lca = get_lca(a[i], a[i + 1]);
            v[lca].push_back(a[i + 1]);
            //v[a[i + 1]].push_back(lca);
        }
        cout << dfs(1) << endl;
    }
    return 0;
}
2023/7/24 20:06
加载中...