Official

D - ボールの箱詰め / Packing Balls into Boxes Editorial by admin

Claude 4.6 Opus (Thinking)

概要

\(K\) 個のボールを「同じ箱に入れなければならない」制約でグループ化し、得られたグループ数 \(G\) 個の区別可能な要素を \(N\) 個の区別可能な箱に全射(どの箱も空でない)で配る場合の数を求める問題です。

考察

制約をグループにまとめる

「ボール \(A\) と \(B\) は同じ箱」「ボール \(B\) と \(C\) は同じ箱」ならば \(A, B, C\) は全て同じ箱に入ります。これは推移的な関係なので、Union-Find(素集合データ構造) を使って、必ず同じ箱に入るボールの集合(グループ)を効率的に求められます。

グループの分配問題に帰着

同じグループのボールは必ず同じ箱に入るため、グループ全体を1つの「塊」として扱えます。グループ数を \(G\) とすると、問題は:

\(G\) 個の区別可能なグループを \(N\) 個の区別可能な箱に分配し、全ての箱に1個以上のグループが入るようにする場合の数

に帰着されます。

全射の数え上げ

これは「\(G\) 要素の集合から \(N\) 要素の集合への全射(surjection)の個数」そのものです。

  • \(G < N\) の場合:グループ数が箱の数より少ないので、必ず空の箱ができてしまい、答えは \(0\) です。
  • \(G \geq N\) の場合:包除原理を使って計算します。

包除原理による全射の公式

\(N\) 個の箱のうち少なくとも1つが空になるケースを排除するため、包除原理を適用します:

\[\text{全射の数} = \sum_{i=0}^{N} (-1)^i \binom{N}{i} (N - i)^G\]

直感的には: - \(i = 0\): 制限なしで \(G\) グループを \(N\) 箱に入れる方法 \(= N^G\) - \(i\) 個の特定の箱を「使わない」と決めたとき、残り \((N-i)\) 箱に入れる方法 \(= (N-i)^G\) - これを \(\binom{N}{i}\) 通りの選び方について足し引きする

アルゴリズム

  1. Union-Find で \(K\) 個のボールをグループ化し、グループ数 \(G\) を求める。
  2. \(G < N\) なら答えは \(0\)。
  3. \(G \geq N\) なら、包除原理の公式 \(\sum_{i=0}^{N} (-1)^i \binom{N}{i} (N-i)^G\) を計算する。
  4. 二項係数の計算には階乗とその逆元を前計算しておく。

具体例

\(N=2, K=3, M=1\) で「ボール1とボール2は同じ箱」の制約がある場合: - グループ: \(\{1, 2\}, \{3\}\) → \(G = 2\) - 全射の数 \(= \sum_{i=0}^{2} (-1)^i \binom{2}{i}(2-i)^2 = 1 \cdot 4 - 2 \cdot 1 + 1 \cdot 0 = 2\) - 分配は「グループ\(\{1,2\}\)→箱1, \(\{3\}\)→箱2」と「グループ\(\{1,2\}\)→箱2, \(\{3\}\)→箱1」の2通り。

計算量

  • 時間計算量: \(O(K \cdot \alpha(K) + N \log G)\)
    • Union-Find の構築: \(O(K \cdot \alpha(K))\)(\(\alpha\) はアッカーマン逆関数)
    • 包除原理のループ: \(O(N \log G)\)(各項で \(\text{pow}(N-i, G, \text{MOD})\) に \(O(\log G)\))
  • 空間計算量: \(O(K + N)\)(Union-Find 配列と階乗テーブル)

実装のポイント

  • Union-Find には経路圧縮とランクによる統合を用いて高速化する。

  • 階乗の逆元は、最大値の逆元から逆順に計算する(フェルマーの小定理で \(\text{fact}[n]^{-1}\) を求め、\(\text{inv\_fact}[i] = \text{inv\_fact}[i+1] \times (i+1)\) で降ろす)。

  • 包除原理で引き算をする際、負の数にならないよう MOD を加えてから余りを取る。

    ソースコード

import sys
from sys import stdin

def main():
    input = stdin.readline
    N, K, M = map(int, stdin.readline().split())
    
    MOD = 998244353
    
    # Union-Find to group balls that must be in the same box
    parent = list(range(K + 1))
    rank = [0] * (K + 1)
    
    def find(x):
        while parent[x] != x:
            parent[x] = parent[parent[x]]
            x = parent[x]
        return x
    
    def union(x, y):
        px, py = find(x), find(y)
        if px == py:
            return
        if rank[px] < rank[py]:
            px, py = py, px
        parent[py] = px
        if rank[px] == rank[py]:
            rank[px] += 1
    
    for _ in range(M):
        u, v = map(int, stdin.readline().split())
        union(u, v)
    
    # Count number of groups
    groups = set()
    for i in range(1, K + 1):
        groups.add(find(i))
    G = len(groups)  # number of groups (each group must go into exactly one box)
    
    # We need to distribute G groups into N boxes such that:
    # - each group goes into exactly one box
    # - no box is empty
    # This is equivalent to surjective functions from G groups to N boxes
    # which is S(G, N) * N! where S is Stirling numbers of the second kind
    # But since groups are distinguishable and boxes are distinguishable,
    # it's the number of surjective functions from a G-element set to an N-element set.
    
    # If G < N, answer is 0 (can't fill all boxes)
    if G < N:
        print(0)
        return
    
    # Number of surjective functions from G-set to N-set:
    # sum_{i=0}^{N} (-1)^i * C(N, i) * (N - i)^G
    # Using inclusion-exclusion
    
    # Precompute factorials and inverse factorials
    max_n = max(N, G) + 1
    fact = [1] * (max_n + 1)
    for i in range(1, max_n + 1):
        fact[i] = fact[i - 1] * i % MOD
    
    inv_fact = [1] * (max_n + 1)
    inv_fact[max_n] = pow(fact[max_n], MOD - 2, MOD)
    for i in range(max_n - 1, -1, -1):
        inv_fact[i] = inv_fact[i + 1] * (i + 1) % MOD
    
    ans = 0
    for i in range(N + 1):
        # C(N, i) * (N - i)^G * (-1)^i
        comb = fact[N] * inv_fact[i] % MOD * inv_fact[N - i] % MOD
        term = comb * pow(N - i, G, MOD) % MOD
        if i % 2 == 0:
            ans = (ans + term) % MOD
        else:
            ans = (ans - term + MOD) % MOD
    
    print(ans % MOD)

main()

この解説は claude4.6opus-thinking によって生成されました。

posted:
last update: