#!/usr/bin/python3import sysimport randomif len(sys.argv) != 4: print(f'Usage: {sys.argv[0]} n q seed') sys.exit(-1)MAXA = 1000MAXC = 10**9n = int(sys.argv[1])q = int(sys.argv[2])random.seed(sys.argv[3])print(n, q)a = (random.randint(1, MAXA) for i in range(n - 1))c = (random.randint(1, MAXC) for i in range(n))print(*a)print(*c)for i in range(q): x, y = (random.randint(1, n) for i in range(2)) print(x, y)
解析
我們認為求解乘法逆元是常數複雜度的。
定義 s(i)=0≤j<i∑ai,即 a 的字首和
初始的 n 三方做法
定義 l(x,y) 表示 x 到 y 的期望距離。首先,當 x=y 時,l(x,y)=0,因為兩點重合。由於 l(x,y) 顯然等於 l(y,x),所以我們可以假設 x<y。因為我們是按照從 0 到 n−1 的順序加的點,所以 y 不可能是 x 的祖先,也就是說 lca(x,y) 一定不等於 y。列舉 y 的所有父親 i,則 l(x,i) 再加上 i 到 y 的邊權就是 l(x,y)。
原來的 l(x,y),就狀態定義就是 O(n2) 種,比較難最佳化,考慮拆分一下。我們知道樹上 x 到 y 的距離等於 x 到根的距離加 y 到根的距離減去 2 倍
lca(x,y) 到根的距離。放到期望也是一樣的,我們定義 d(x) 表示
x 到根的期望距離,f(x,y) 表示 lca(x,y) 到根的距離,則
l(x,y)=d(x)+d(y)−2f(x,y)。這樣拆分降低了耦合度。
求解 d(x)
和求 l(x,y) 的方法類似,列舉 x 的父親 i,然後把 d(i) 加上 i 到 x
的邊權。容易得到:
d(x)=⎩⎨⎧0s(i)0≤i<x∑ai×[d(i)+ci]+cxx=0其他
可以維護 0≤i<x∑ai×[d(i)+ci] 做到 O(n) 求解。
求解 f(x, y)
仍然和求 l(x,y) 的方法類似,仍然是假設 x<y,列舉 y 的父親 i,由於
y 不是 x 的祖先,所以如果 y 的父親為 i,則 lca 到根距離的期望就是 f(x,i)。容易得到:
using mint = static_modint<1'000'000'007>;auto get_d(const auto &a, const auto &sa, const auto &c){ const int n = a.size(); std::vector<mint> d(n); mint sd = 0; d[0] = 0; sd += mint{c[0]} * a[0]; for (int i = 1; i < n; i++) { d[i] = sd / sa[i] + c[i]; sd += (d[i] + c[i]) * a[i]; } return d;}// g(i) = f(i, i + 1)auto get_g(const auto &a, const auto &sa, const auto &d){ const int n = a.size(); std::vector<mint> g(n); mint s = 0; for (int i = 0; i + 1 < n; i++) { g[i] = (d[i] * a[i] + s) / sa[i + 1]; s += g[i] * a[i]; } return g;};int main(){ int n, q; scanf("%d%d", &n, &q); std::vector<int> a(n), c(n); for (int i = 0; i < n - 1; i++) scanf("%d", &a[i]); for (auto &i : c) scanf("%d", &i); std::vector<mint> sa(n); for (int i = 0; i < n - 1; i++) sa[i + 1] = sa[i] + a[i]; const auto d = get_d(a, sa, c); const auto g = get_g(a, sa, d); for (int i = 0; i < q; i++) { int x, y; scanf("%d%d", &x, &y); x--; y--; if (x == y) { puts("0"); } else { int t = std::min(x, y); printf("%d\n", (d[x] + d[y] - g[t] * 2).val()); } }}