AcWing 885. 求组合数 I
本节目标
预处理 Pascal 三角回答小范围组合数查询。
题意与约束
给定多组 a, b,输出 C(a,b) 对 1_000_000_007 取模后的值。本题的 a 上界较小,适合先建整张 Pascal 三角。
第一反应与瓶颈
递归使用 C(a,b)=C(a-1,b-1)+C(a-1,b) 会反复求同一子问题;逐组建表也重复了前缀行。
数学关系与算法推导
第 a 行的两端 C(a,0)、C(a,a) 都是 1,中间每格由左上和右上相加。先读取全部查询确定最大 a,再一次性预处理到该行。
正确性依据
任意大小为 b 的选择集要么包含第 a 个元素,要么不包含它,两类分别对应 C(a-1,b-1) 和 C(a-1,b),且互不重叠并覆盖全部情况。
样例执行过程
C(5,3)=C(4,2)+C(4,3)=6+4=10;C(5,0) 与 C(5,5) 都直接落在三角形边界上。
代码实现
- C++
- Python
C++17
#include <algorithm>
#include <iostream>
#include <utility>
#include <vector>
using namespace std;
constexpr int MOD = 1000000007;
vector<int> combinationQueriesPascal(const vector<pair<int, int>>& queries) {
int maximum = 0;
for (const auto& [a, b] : queries) {
maximum = max(maximum, a);
}
vector<vector<int>> combinations(maximum + 1, vector<int>(maximum + 1, 0));
combinations[0][0] = 1;
for (int a = 1; a <= maximum; a++) {
combinations[a][0] = 1;
combinations[a][a] = 1;
for (int b = 1; b < a; b++) {
combinations[a][b] = (combinations[a - 1][b - 1] + combinations[a - 1][b]) % MOD;
}
}
vector<int> answers;
answers.reserve(queries.size());
for (const auto& [a, b] : queries) {
answers.push_back(b < 0 || b > a ? 0 : combinations[a][b]);
}
return answers;
}
#ifndef ALGORITHM_TUTORIAL_NO_MAIN
int main() {
ios::sync_with_stdio(false);
cin.tie(nullptr);
int count;
cin >> count;
vector<pair<int, int>> queries(count);
for (auto& [a, b] : queries) {
cin >> a >> b;
}
for (int answer : combinationQueriesPascal(queries)) {
cout << answer << '\n';
}
return 0;
}
#endif
Python 3
import sys
MOD = 1_000_000_007
def combination_queries_pascal(queries: list[tuple[int, int]]) -> list[int]:
maximum = max((a for a, _ in queries), default=0)
combinations = [[0] * (maximum + 1) for _ in range(maximum + 1)]
combinations[0][0] = 1
for a in range(1, maximum + 1):
combinations[a][0] = 1
combinations[a][a] = 1
for b in range(1, a):
combinations[a][b] = (combinations[a - 1][b - 1] + combinations[a - 1][b]) % MOD
return [0 if b < 0 or b > a else combinations[a][b] for a, b in queries]
def main() -> None:
data = list(map(int, sys.stdin.buffer.read().split()))
queries = [(data[index], data[index + 1]) for index in range(1, len(data), 2)]
print('\n'.join(map(str, combination_queries_pascal(queries))))
if __name__ == '__main__':
main()
复杂度分析
设最大行号为 N,预处理时间和空间均为 O(N²);每个查询 O(1)。
边界与易错点
C(0,0)=1。两侧边界必须先设为 1;若输入的 b 不在 0..a,核心函数返回 0,而合法平台输入不会触发这一分支。
模式迁移
当 N 很大但查询仍在质数模数下时,改用阶乘与逆阶乘预处理。回到组合计数。