Expected Median 的题解


记住只在没有思路时使用题解,不要从它复制粘贴代码。请尊重题目和题解的作者。
在解题之前提交题解的代码会导致封禁。

作者: admin

概述

本题要求对二进制数组 a 的所有长度为奇数 k 的子序列,计算其中位数为 1 的子序列个数之和。核心解法是将中位数为 1 的条件转化为子序列中 1 的数量至少为 \frac{k+1}{2},然后用组合数枚举并求和。

分析
核心观察

长度为奇数 k 的子序列,排序后中位数为第 \frac{k+1}{2} 个元素。由于元素仅为 0 或 1,中位数为 1 当且仅当子序列中 1 的个数严格大于 0 的个数,即至少包含 \lfloor \frac{k}{2} \rfloor + 1 个 1。

思路

设整个数组中有 x 个 1,y 个 0(显然 y = n - x)。枚举子序列中 1 的数量 i,则 0 的数量为 k - i。为了使中位数为 1,需要满足 i \ge \lfloor \frac{k}{2} \rfloor + 1,同时 0 \le i \le x 且 0 \le k - i \le y。对于每个合法的 i,选择 i 个 1 的方式有 \binom{x}{i} 种,选择 k-i 个 0 的方式有 \binom{y}{k-i} 种。因此总贡献为:

\displaystyle  \sum_{i = \lfloor k/2 \rfloor + 1}^{\min(x, k)} \binom{x}{i} \binom{y}{k-i}

将所有测试用例的答案累加即可。组合数通过预处理阶乘和逆元在 O(1) 时间内计算,整体时间复杂度为 O(\sum k),由于所有 n 之和不超过 2 \times 10^5,可以接受。

具体示例

以样例第一组 n=4, k=3, a=[1,0,0,1]:x=2, y=2,下限 \lfloor 3/2 \rfloor + 1 = 2,枚举 i=2(i=3 不可能因为 x=2):C(2,2)*C(2,1)=1*2=2,答案 2。
第二组 n=5, k=1, a=[1,...]:x=5, y=0,下限 \lfloor 1/2 \rfloor + 1 = 1,枚举 i=1:C(5,1)*C(0,0)=5,答案 5。

算法步骤
  1. 预处理阶乘数组 fact,长度至最大可能的 n(2 \times 10^5),并利用费马小定理预处理逆元(或在组合数中实时用快速幂求逆)。
  2. 对每个测试用例,读入 n, k 和数组 a,统计其中 1 的个数 ones,零的个数 zeros = n - ones。
  3. 确定枚举下界 need = k / 2 + 1(因为 k 为奇数)。
  4. 初始化答案 ans = 0。
  5. 对 i 从 need 到 min(ones, k) 遍历:
    • 计算 ways = C(ones, i) * C(zeros, k - i) % MOD。
    • 将 ways 累加到 ans 中,并取模。
  6. 输出 ans。
复杂度分析
  • 时间:预计算阶乘 O(N),其中 N = 2 \times 10^5;每个测试用例枚举 O(\min(ones, k) - need + 1),总枚举次数不超过所有测试用例的 n 之和,因此总时间复杂度为 O(N + \sum n)。
  • 空间:O(N) 用于阶乘数组。
实现注意事项
  • 组合数 C(n, k) 在 n < k 时返回 0。
  • 使用 long long 类型存储阶乘和中间结果,避免乘法溢出(模数 10^9+7 在 int 范围内,但乘积需用 long long)。
  • 逆元可用快速幂计算 fact[n-k] * fact[k] % MOD 的 MOD-2 次幂,也可预先计算逆元数组。参考代码使用快速幂,每次组合数调用一次快速幂,但枚举次数总和不超 2e5,效率足够。
  • 注意所有测试用例的 n 之和不超过 2 \times 10^5,因此阶乘预计算到最大值即可。
源代码
#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()

评论

目前没有评论。