刚刚比赛D部分分求助
  • 板块学术版
  • 楼主Erine
  • 当前回复0
  • 已保存回复0
  • 发布时间2023/1/27 18:50
  • 上次更新2023/10/24 02:55:54
查看原帖
刚刚比赛D部分分求助
738474
Erine楼主2023/1/27 18:50

打的65pts。只有15pts

#include <bits/stdc++.h>
#define int long long

using namespace std;

const int maxn = 5e5 + 1;
const int mod = 1e9 + 7;
const int maxd = 20;

struct edge {
    int to, next;
};

int n, q;
int head[maxn];
edge e[maxn << 1];
int cnt;

void add_edge(int u, int v) {
    e[++cnt].to = v;
    e[cnt].next = head[u];
    head[u] = cnt;
}

int fac[maxn];

void init() {
    fac[0] = 1;
    for (int i = 1; i < maxn; i++) fac[i] = fac[i - 1] * i % mod;
}

int fa[maxn][maxd];
int dep[maxn];

void get_dep(int u, int depth, int fat) {
    dep[u] = depth;
    for (int i = head[u]; i; i = e[i].next) {
        if (e[i].to != fat) {
        	get_dep(e[i].to, depth + 1, u);
		}
    }
} 

void get_fa(int cur, int fat) {
    fa[cur][0] = fat;
    for (int i = 1; i <= log2(dep[cur]) + 1; i++) {
        fa[cur][i] = fa[fa[cur][i - 1]][i - 1];
    }
    for (int i = head[cur]; i; i = e[i].next) {
        if (e[i].to != fat) {
        	get_fa(e[i].to, cur);
		}
    }
}

int lca(int u, int v) {
    if (dep[u] < dep[v]) swap(u, v);
    while (dep[u] > dep[v]) {
        u = fa[u][(int) log2(dep[u] - dep[v])];
    }
    if (u == v) return u;
    for (int i = log2(dep[u]); i >= 0; i--) {
        if (fa[u][i] != fa[v][i]) {
            u = fa[u][i];
            v = fa[v][i];
        }
    }
    return fa[u][0];
}

int dist(int u, int v) {
    return dep[u] + dep[v] - 2 * dep[lca(u, v)];
}

void dfs(int u, int fa, int ub, int vb, int &mul) {
    int tot = 0;
    for (int i = head[u]; i; i = e[i].next) {
        int v = e[i].to;
        if (v == fa || (u == ub && v == vb) || (u == vb && v == ub)) continue;
        tot++;
        dfs(v, u, ub, vb, mul);
    }
    mul = mul * fac[tot] % mod;
}

signed main() {
    init();
    cin >> n >> q;
    for (int i = 1; i < n; i++) {
        int u, v;
        cin >> u >> v;
        add_edge(u, v);
        add_edge(v, u);
    }
    get_dep(1, 0, 0);
    get_fa(1, 0);
    int tot = 1;
    dfs(1, 0, 0, 0, tot);
    if (tot == 1) {
        assert(false);
        while (q--) {
            int x, y;
            cin >> x >> y;
            if (dist(x, y) == 1) {
                cout << 1 << endl;
            } else cout << 2 << endl;
        }
        return 0;
    }
    while (q--) {
        int x, y;
        cin >> x >> y;
        int z = lca(x, y);
        if (dist(x, y) == 1) {
            cout << tot << endl;
            continue;
        }
        int xt, yt;
        for (int i = head[z]; i; i = e[i].next) {
            int v = e[i].to;
            if (v == fa[z][0]) continue;
            if (lca(x, v) == v) xt = v;
            if (lca(y, v) == v) yt = v;
        }
        int ans1 = 1, ans2 = 1;
        dfs(1, 0, z, xt, ans1);
        dfs(1, 0, z, yt, ans2);
        cout << (ans1 + ans2) % mod << endl;
    }
    return 0;
}
2023/1/27 18:50
加载中...