题解:CF961G Partitions

· · 题解

一种力大砖飞的做法。

考虑拆贡献:对于一个 w_i,枚举其所属子集大小 s,则需要从剩下的 n-1 个元素中再挑选 s-1 个,与这个 w_i 合成一个大小为 s 的子集,根据定义有 \displaystyle\binom{n-1}{s-1} 种方案。

同时,我们需要把剩下的 n-s 个元素划分为恰好 k-1 个非空子集,根据第二类斯特林数的定义,有 \begin{Bmatrix}n-s\\k-1\end{Bmatrix} 种方案。

根据乘法原理,w_i 属于大小为 s 的子集的划分方案数即为 \displaystyle\binom{n-1}{s-1}\begin{Bmatrix}n-s\\k-1\end{Bmatrix}。每种方案中 w_i 的贡献均为 s\cdot w_i,于是可以写出答案:

\begin{aligned} \mathrm{answer}&=\sum_{i=1}^n\sum_{s=1}^{n-(k-1)}s\cdot w_i\binom{n-1}{s-1}\begin{Bmatrix}n-s\\k-1\end{Bmatrix}&(1)\\ &=\left(\sum_{i=1}^nw_i\right)\left(\sum_{s=1}^{n-k+1}s\binom{n-1}{s-1}\begin{Bmatrix}n-s\\k-1\end{Bmatrix}\right)&(2) \end{aligned}
  1. 解释:枚举每个 w_i 以及所有可能的子集大小 s,将其贡献乘上方案数并求和。

  2. 将与内层循环无关的项 w_i 依据乘法分配律提前。

这个式子还可以简化,但是已经可以直接力大砖飞解掉了。组合数可以预处理阶乘并 \mathcal O(1) 求出,难点在于后面的第二类斯特林数。

我们注意到这里斯特林数的 k-1 不随求和改变,也就是说我们要求的所有斯特林数位于同一列,这可通过生成函数 \mathcal O(n\log n) 求出。注意到此处的模数为 10^9+7 不是一个 NTT 模数,所以需要使用任意模数多项式乘法来实现第二类斯特林数·列。

:::info[code]

#include <bits/extc++.h>

namespace {
    constexpr unsigned Modulus = 1E9 + 7;
    std::vector<std::complex<double>> factors;

    auto DIF(std::vector<std::complex<double>>& A) -> void {
        for (auto n = unsigned(A.size()), m = n >> 1; m != 0; m >>= 1) {
            for (auto s = m << 1, o = n / s, i = 0U; i != n; i += s) {
                for (auto j = 0U; j != m; ++j) {
                    auto x = A[i + j], y = A[i + j + m];
                    A[i + j] = x + y, A[i + j + m] = (x - y) * factors[j * o];
                }
            }
        }
    }

    auto DIT(std::vector<std::complex<double>>& A) -> void {
        for (auto n = unsigned(A.size()), m = 1U; m != n; m <<= 1) {
            for (auto s = m << 1, o = n / s, i = 0U; i != n; i += s) {
                for (auto j = 0U; j != m; ++j) {
                    auto x = A[i + j], y = A[i + j + m] * std::conj(factors[j * o]);
                    A[i + j] = x + y, A[i + j + m] = x - y;
                }
            }
        }
    }

    auto multiplies(std::vector<unsigned> P, std::vector<unsigned> Q) -> std::vector<unsigned> {
        auto s = P.size() + Q.size();
        --s;
        auto n = unsigned(std::__bit_ceil(s));
        auto m = unsigned(std::__bit_width(n) - 1);
        factors.resize(n >> 1);
        for (auto i = 0U; i != n >> 1; ++i) {
            auto Pi = std::acos(-1);
            factors[i] = std::polar(1.0, -2.0 * Pi * i / n);
        }
        std::vector<unsigned> reversed(n);
        for (auto i = 1U; i != n; ++i) {
            reversed[i] = reversed[i >> 1] >> 1;
            reversed[i] |= (i & 1) << (m - 1);
        }
        std::vector<std::complex<double>> A(n), B(n), C(n), D(n);
        for (auto i = 0U; i != P.size(); ++i) {
            auto x = P[i] % Modulus;
            A[i] = {double(x & 0xFFFF), double(x >> 16)};
        }
        for (auto i = 0U; i != Q.size(); ++i) {
            auto x = Q[i] % Modulus;
            B[i] = {double(x & 0xFFFF), double(x >> 16)};
        }
        DIF(A);
        DIF(B);
        for (auto o = 0U; o != n; ++o) {
            auto i = reversed[o];
            auto j = reversed[(n - o) & (n - 1)];
            auto u = (A[i] + std::conj(A[j])) * 0.5;
            auto v = (A[i] - std::conj(A[j])) * std::complex(0.0, -0.5);
            auto x = (B[i] + std::conj(B[j])) * 0.5;
            auto y = (B[i] - std::conj(B[j])) * std::complex(0.0, -0.5);
            C[i] = u * x + std::complex<double>(0.0, 1.0) * v * y;
            D[i] = u * y + v * x;
        }
        DIT(C);
        DIT(D);
        std::vector<unsigned> result(s);
        for (auto i = 0U; i != s; ++i) {
            auto u = std::uint64_t(std::int64_t(std::llround(C[i].real() / n)) % Modulus + Modulus) % Modulus;
            auto v = std::uint64_t(std::int64_t(std::llround(D[i].real() / n)) % Modulus + Modulus) % Modulus;
            auto w = std::uint64_t(std::int64_t(std::llround(C[i].imag() / n)) % Modulus + Modulus) % Modulus;
            result[i] = unsigned((u + (v << 16) + (w << 32)) % Modulus);
        }
        return result;
    }
}

auto main() -> int {
    std::cin.tie(nullptr)->sync_with_stdio(false);

    unsigned n, k;
    std::cin >> n >> k;

    std::vector<unsigned> W(n);
    std::copy_n(std::istream_iterator<unsigned>(std::cin), n, W.data());

    std::vector<unsigned> stirling(n);
    const auto required = n - --k;

    auto product = [required](auto& self, auto l, auto r) -> std::vector<unsigned> {
        switch (auto s = r - l) {
            case 0: return {1};
            case 1: return {1, Modulus - l};
            default: {
                const auto m = (l + r) / 2;
                auto result = multiplies(self(self, l, m), self(self, m, r));
                if (result.size() > required) result.resize(required);
                return result;
            }
        }
    };

    const auto denominator = product(product, 1U, k + 1);
    std::vector<unsigned> inverse{1};
    while (inverse.size() < required) {
        const auto size = std::min<std::size_t>(inverse.size() * 2, required);
        std::vector prefix(denominator.data(), denominator.data() + std::min(denominator.size(), size));
        auto correction = multiplies(prefix, inverse);
        correction.resize(size);
        for (auto& x : correction) x = x == 0 ? 0 : Modulus - x;
        correction[0] = (correction[0] + 2) % Modulus;
        (inverse = multiplies(inverse, correction)).resize(size);
    }
    std::copy(inverse.begin(), inverse.end(), stirling.begin() + k);

    std::vector factorial(n, 0ULL);
    for (auto i = factorial[0] = 1; i != n; ++i)
        factorial[i] = factorial[i - 1] * i % Modulus;

    auto power = [](auto v) {
        auto o = 1ULL;
        for (auto e = Modulus - 2; e; (e & 1) && (o = o * v % Modulus), v = v * v % Modulus, e >>= 1);
        return o;
    };

    auto sum = 0ULL;
    for (auto s = 1ULL; s <= n - k; ++s) {
        auto binomial = factorial.back() % Modulus * power(factorial[s - 1]) % Modulus * power(factorial[n - s]) % Modulus;
        sum = (sum + s * stirling[n - s] % Modulus * binomial % Modulus) % Modulus;
    }
    sum *= std::accumulate(W.begin(), W.end(), 0ULL) % Modulus;
    std::cout << sum % Modulus << '\n';

    constexpr auto Anon = 0x0908, Soyo = 0x0527, Tomori = 0x1122, Taki = 0x0809, Rana = 0x0222;
    std::clog << "Duration = " << std::clock() << " clocks\n";
    return Anon & Soyo & Tomori & Taki & Rana;
}

:::