Expected Median 的题解
记住只在没有思路时使用题解,不要从它复制粘贴代码。请尊重题目和题解的作者。
在解题之前提交题解的代码会导致封禁。
在解题之前提交题解的代码会导致封禁。
作者:
概述
本题要求对二进制数组 的所有长度为奇数
的子序列,计算其中位数为
的子序列个数之和。核心解法是将中位数为
的条件转化为子序列中
的数量至少为
,然后用组合数枚举并求和。
分析
核心观察
长度为奇数 的子序列,排序后中位数为第
个元素。由于元素仅为
或
,中位数为
当且仅当子序列中
的个数严格大于
的个数,即至少包含
个
。
思路
设整个数组中有 个
,
个
(显然
)。枚举子序列中
的数量
,则
的数量为
。为了使中位数为
,需要满足
,同时
且
。对于每个合法的
,选择
个
的方式有
种,选择
个
的方式有
种。因此总贡献为:
将所有测试用例的答案累加即可。组合数通过预处理阶乘和逆元在 时间内计算,整体时间复杂度为
,由于所有
之和不超过
,可以接受。
具体示例
以样例第一组 ,
,
:
,下限
,枚举
(
不可能因为
):
,答案
。
第二组 ,
,
:
,下限
,枚举
:
,答案
。
算法步骤
- 预处理阶乘数组
fact,长度至最大可能的(
),并利用费马小定理预处理逆元(或在组合数中实时用快速幂求逆)。
- 对每个测试用例,读入
和数组
,统计其中
的个数
ones,零的个数zeros = n - ones。 - 确定枚举下界
need = k / 2 + 1(因为为奇数)。
- 初始化答案
ans = 0。 - 对
i从need到min(ones, k)遍历:- 计算
ways = C(ones, i) * C(zeros, k - i) % MOD。 - 将
ways累加到ans中,并取模。
- 计算
- 输出
ans。
复杂度分析
- 时间:预计算阶乘
,其中
;每个测试用例枚举
,总枚举次数不超过所有测试用例的
之和,因此总时间复杂度为
。
- 空间:
用于阶乘数组。
实现注意事项
- 组合数
C(n, k)在n < k时返回。
- 使用
long long类型存储阶乘和中间结果,避免乘法溢出(模数在
int范围内,但乘积需用long long)。 - 逆元可用快速幂计算
fact[n-k] * fact[k] % MOD的次幂,也可预先计算逆元数组。参考代码使用快速幂,每次组合数调用一次快速幂,但枚举次数总和不超
,效率足够。
- 注意所有测试用例的
之和不超过
,因此阶乘预计算到最大值即可。
源代码
#include <bits/stdc++.h>
using namespace std;
const int MAXN = 200000 + 5;
const int MOD = 1000000007;
long long fact[MAXN];
long long mod_pow(long long a, long long b) {
long long res = 1;
while (b > 0) {
if (b & 1) res = res * a % MOD;
a = a * a % MOD;
b >>= 1;
}
return res;
}
long long C(int n, int k) {
if (n < k || k < 0) return 0;
return fact[n] * mod_pow(fact[n - k] * fact[k] % MOD, MOD - 2) % MOD;
}
int main() {
ios::sync_with_stdio(false);
cin.tie(nullptr);
fact[0] = 1;
for (int i = 1; i < MAXN; ++i) {
fact[i] = fact[i - 1] * i % MOD;
}
int t;
cin >> t;
while (t--) {
int n, k;
cin >> n >> k;
int ones = 0;
for (int i = 0; i < n; ++i) {
int x;
cin >> x;
ones += x;
}
int zeros = n - ones;
int need = k / 2 + 1;
long long ans = 0;
int upper = min(ones, k);
for (int i = need; i <= upper; ++i) {
long long ways = C(ones, i) * C(zeros, k - i) % MOD;
ans += ways;
if (ans >= MOD) ans -= MOD;
}
cout << ans << '\n';
}
return 0;
}
import sys
MOD = 10**9 + 7
MAXN = 200000 + 5
fact = [1] * MAXN
for i in range(1, MAXN):
fact[i] = fact[i-1] * i % MOD
def mod_pow(a, b):
res = 1
while b:
if b & 1:
res = res * a % MOD
a = a * a % MOD
b >>= 1
return res
def C(n, k):
if n < k or k < 0:
return 0
return fact[n] * mod_pow(fact[n-k] * fact[k] % MOD, MOD - 2) % MOD
def solve():
data = sys.stdin.buffer.read().split()
t = int(data[0])
idx = 1
out = []
for _ in range(t):
n = int(data[idx]); k = int(data[idx+1]); idx += 2
arr = list(map(int, data[idx:idx+n])); idx += n
ones = sum(arr)
zeros = n - ones
need = k // 2 + 1
ans = 0
upper = min(ones, k)
for i in range(need, upper + 1):
ways = C(ones, i) * C(zeros, k - i) % MOD
ans += ways
if ans >= MOD:
ans -= MOD
out.append(str(ans))
sys.stdout.write("\n".join(out))
if __name__ == "__main__":
solve()
评论