CF题解——Fenwick Tree

C. Fenwick Tree 解题思路

核心问题分析

题意:定义 lowbit(x)\text{lowbit}(x)xx 二进制最低位的值。对长度为 nn 的数组 aa,定义它的树状数组变换

f(a)k=(i=klowbit(k)+1kai)mod998244353f(a)_k = \Big(\sum_{i = k - \text{lowbit}(k) + 1}^{k} a_i\Big) \bmod 998244353

再定义 fkf^k 为把 ff 迭代作用 kk 次。现在给你 bbkk,要求还原出一个 aa 使得 fk(a)=bf^k(a) = b

n2×105n \le 2 \times 10^5,但是 k109k \le 10^9——这个 kk 大得离谱,明摆着不可能真的迭代 kk 次。多组数据,n2×105\sum n \le 2\times 10^5

1. 把树状数组还原成一棵树

第一步是把 ff 这个看起来很"位运算"的东西还原成它的本来面目。众所周知,树状数组的下标之间存在一个天然的父子关系:

parent(i)=i+lowbit(i)\text{parent}(i) = i + \text{lowbit}(i)

在这个结构下,sk=i=klowbit(k)+1kais_k = \sum_{i = k - \text{lowbit}(k)+1}^{k} a_i 恰好就是「kk 的子树内所有 aia_i 之和」(包含 kk 自己)。

所以 ff 用矩阵语言写出来就是

f=I+Uf = I + U

其中 Uk,i=1U_{k,i} = 1 当且仅当 iikk真后代。这是一个幂零(nilpotent)的严格下三角结构——树的深度有限,UU 乘若干次就变成零矩阵。

2. 关键性质:树上路径唯一,所以 UjU^j 长得极其干净

一般矩阵的幂会很难看,但这里有个救命的性质:每个点只有一个父亲,所以从 kk 往下走 jj 步到达某个具体节点 xx 的路径,要么不存在,要么恰好只有一条。于是

(Uj)k,x=[x 是 k 的后代,且深度恰好差 j](U^j)_{k,x} = [\,x \text{ 是 } k \text{ 的后代,且深度恰好差 } j\,]

也就是说 UjU^j 依然只是一个 01 矩阵,没有任何"路径计数"的爆炸。这就是这道题能做的根本原因。

3. 不求 fkf^k,直接求 fkf^{-k}

quot;">​

我们要的是 a=fk(b)a = f^{-k}(b),所以与其去算 fkf^k 再解方程,不如直接把 fkf^{-k} 求出来。

先看 f1f^{-1} 长什么样。sks_k 是子树和,那么把 sks_k 减去它所有直接儿子的子树和,剩下的就是 aka_k 自己:

f1=IV,Vk,c=[c 是 k 的直接儿子]f^{-1} = I - V, \qquad V_{k,c} = [\,c \text{ 是 } k \text{ 的直接儿子}\,]

于是用二项式定理(IIVV 显然可交换):

fk=(IV)k=j0(1)j(kj)Vjf^{-k} = (I - V)^k = \sum_{j \ge 0} (-1)^j \binom{k}{j} V^j

而根据上一节的路径唯一性,VjV^j 就是"深度恰好差 jj 的后代"关系。翻译成人话:

a_x = \sum_{j \ge 0} (-1)^j \dbinom{k}{j} \cdot \big(\text{x$ 的所有深度为 } j \text{ 的后代的 } b \text{ 值之和}\big)$

kk 高达 10910^9 完全不是问题——它只出现在组合数 (kj)\binom{k}{j} 里,而 jj 的范围很小。

4. jj 到底能有多大?

沿着 ii+lowbit(i)i \to i + \text{lowbit}(i) 往上走,每走一步 lowbit\text{lowbit} 至少翻倍,所以从任意节点走到根最多只需要 O(logn)O(\log n) 步。n2×105n \le 2\times10^5 时这个上界不超过 1818,代码里开到 2020 稳妥有余。

组合数 (kj)\binom{k}{j} 也不用预处理阶乘——jj 只有 2020 个取值,用递推一路乘过去就行:

(kj)=(kj1)kj+1j\binom{k}{j} = \binom{k}{j-1} \cdot \frac{k-j+1}{j}

cpp
int cur_C = 1;
cal[0] = 1;
for (int j = 1; j < max_j; j++) {
    if (k < j) break;                       // j > k 时组合数为 0,后面全是 0
    cur_C = cur_C * ((k - j + 1) % MOD) % MOD * inv(j) % MOD;
    cal[j] = (j & 1) ? (MOD - cur_C) % MOD : cur_C;   // 带上 (-1)^j
}

注意 kk 本身要先取模,但 k<jk < j 的判断必须用取模前的原值——否则 kk 很大时会误判。

5. 换个方向扫:与其"往下找后代",不如"往上找祖先"

公式写的是"对每个 xx,找它深度为 jj 的所有后代"。但直接实现这件事需要建出整棵树再做 DFS,比较麻烦。

反过来想:从每个 ii 出发往上走它的祖先链,走 jj 步到达的那个祖先 xx,恰好就把 bib_i 贡献给了 axa_x。因为祖先链是唯一的,这样扫一遍就不重不漏地覆盖了所有 (x,i,j)(x, i, j) 三元组:

cpp
for (int i = 1; i <= n; i++) {
    int curr = i, j = 0;
    while (curr <= n && j < max_j) {
        if (!cal[j]) break;
        res[curr] = (res[curr] + a[i] * cal[j]) % MOD;
        curr += low_bit(curr);   // 往上跳一层
        j++;
    }
}

一行 curr += low_bit(curr) 就搞定了建树、找祖先、算深度三件事——这道题从头到尾根本不需要显式建出任何树结构,全靠 lowbit\text{lowbit} 的代数性质。

6. 复杂度

外层枚举 nn 个位置,每个位置往上跳 O(logn)O(\log n) 步,组合数预处理是 O(lognlogP)O(\log n \log P) 的常数级别。总复杂度

O(nlogn)O(n \log n)

n2×105\sum n \le 2\times10^533 秒时限下随便跑。

回头看这道 2300*2300 的题,题面把 lowbit\text{lowbit} 摆在最显眼的位置,反而容易让人一头栽进位运算的细节里出不来。真正的破题点是那句"树状数组本质是一棵树,ff 本质是 I+UI + U"——一旦把它写成矩阵,kk 次幂就变成了二项式展开,10910^9 这个吓人的数字瞬间蒸发成 2020 个组合数系数。而让整件事成立的技术前提,是树上父亲唯一导致的 UjU^j 不会退化成路径计数。这三层(树结构 → 线性算子 → 二项式反演)串起来,题就通了。

CPP 代码实现

cpp
// C. Fenwick Tree

#include <bits/stdc++.h>
#define lg(x) (63 - __builtin_clzll(x))
#define all(x) (x).begin(), (x).end()
#define low_bit(x) ((x) & (-x))
#define pb push_back
#define db long double
#define int long long
#define sz(x) (int)x.size()
#define endl "\n"

using namespace std;

const int MOD = 998244353;

void solve() {

    int n, k;
    cin >> n >> k;
    vector<int> a(n + 1);
    for (int i = 1; i <= n; i++) cin >> a[i];

    auto qpow = [&](int base, int exp) {
        int res = 1;
        base %= MOD;
        while (exp > 0) {
            if (exp & 1) res = res * base % MOD;
            base = base * base % MOD;
            exp >>= 1;
        }
        return res;
    };
    auto inv = [&](int x) { return qpow(x, MOD - 2); };

    // cal[j] = (-1)^j * C(k, j),j 最多到树高 O(log n)
    int max_j = 20;
    vector<int> cal(max_j, 0);
    cal[0] = 1;
    int cur_C = 1;
    for (int j = 1; j < max_j; j++) {
        if (k < j) break;
        cur_C = cur_C * ((k - j + 1) % MOD) % MOD * inv(j) % MOD;
        cal[j] = (j & 1) ? (MOD - cur_C) % MOD : cur_C;
    }

    vector<int> res(n + 1, 0);
    for (int i = 1; i <= n; i++) {
        int curr = i, j = 0;
        while (curr <= n && j < max_j) {
            if (!cal[j]) break;
            res[curr] = (res[curr] + a[i] * cal[j]) % MOD;
            curr += low_bit(curr);
            j++;
        }
    }

    for (int i = 1; i <= n; i++) cout << res[i] << " ";
    cout << endl;

}

signed main() {

    ios_base::sync_with_stdio(false);
    cin.tie(nullptr);

    int t = 1;
    cin >> t;

    while (t--) {
        solve();
    }

}
CF题解——Graph Cutting
CF题解——Xor-MST