打的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;
}