Official

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

gemini-3.5-flash-thinking

概要

この問題は、ボールの間に「同じ箱に入れなければならない」という制約がある中で、すべての箱に1個以上のボールが入るように \(K\) 個のボールを \(N\) 個の箱に分配する方法の総数を求める問題です。

一見複雑な条件ですが、Union-Find を用いたグループ化と、包除原理 を用いた重複のない数え上げを組み合わせることで、効率的に解くことができます。


考察

1. ボールのグループ化(連結成分の抽出)

「ボール \(U_j\) とボール \(V_j\) は同じ箱に入れなければならない」という制約は、ボール同士のつながりを表しています。この関係は推移的(AとB、BとCが同じなら、AとCも同じ)であるため、ボール全体をいくつかの「同じ箱に入れなければならないグループ」に分割できます。

これは、ボールを頂点、制約を辺とした無向グラフにおいて、連結成分を求めることと同値です。 ボールの数を \(K\) とし、制約をすべて適用したあとのグループの総数を \(G\) とします。このグループ分けは Union-Find(素集合データ構造) を用いることで効率的に行えます。

2. 箱への分配問題への帰着

同じグループに属するボールはすべて同一の箱に入れる必要があるため、これ以降は「ボール」ではなく「グループ」を単位として考えます。 各グループはそれぞれ区別され、箱も区別されます。 したがって、問題は以下のようにシンプルに言い換えられます。

「\(G\) 個の区別できるグループを、空箱ができないように \(N\) 個の区別できる箱に分配する方法は何通りか?」

まず、グループ数 \(G\) が箱の数 \(N\) より小さい場合(\(G < N\))、どうしても空の箱ができてしまうため、分配方法は \(0\) 通りとなります。

3. 空箱がない分配数の計算(包除原理)

\(G \geq N\) のとき、空箱がないように配る方法を考えます。 もし「空箱があってもよい」という条件であれば、各グループは \(N\) 個の箱のどれに入ってもよいため、配り方は \(N^G\) 通りです。 しかし、今回は「空箱があってはならない」という制約があるため、包除原理(Inclusion-Exclusion Principle) を使用して、空箱がある場合を排除します。

空の箱の個数に着目して、全体から引いていきます。 - すべての配り方(制限なし): \(N^G\) 通り - 少なくとも \(1\) つの特定の箱が空である配り方: \(\binom{N}{1} (N-1)^G\) 通り - 少なくとも \(2\) つの特定の箱が空である配り方: \(\binom{N}{2} (N-2)^G\) 通り - \(\dots\) - 少なくとも \(i\) つの特定の箱が空である配り方: \(\binom{N}{i} (N-i)^G\) 通り

包除原理により、空箱が \(0\) 個である(すべての箱に1つ以上のグループが入る)配り方の総数は、以下の式で計算できます。

\[ \sum_{i=0}^{N} (-1)^i \binom{N}{i} (N-i)^G \]

この式を \(i = 0\) から \(N\) までループを回して計算することで、答えを求めることができます。


アルゴリズム

  1. Union-Findによるグループ化: \(K\) 個の要素を持つ Union-Find を用意し、与えられた \(M\) 個の制約 \((U_j, V_j)\) について union 操作を行います。 最終的なグループ数 \(G\) は、Union-Find 内の代表元の個数(連結成分数)となります。
  2. 判定: \(G < N\) であれば、空箱をなくすことが不可能なため 0 を出力して終了します。
  3. 階乗・逆元の前処理: 包除原理の式に含まれる二項係数 \(\binom{N}{i} = \frac{N!}{i!(N-i)!}\) を高速に計算するため、あらかじめ階乗 \(N!\) とその逆元 \((N!)^{-1} \pmod{998244353}\) の配列を \(O(N)\) で作成しておきます。
  4. 包除原理の計算: \(i = 0\) から \(N\) までループを回し、各項を足し引きして答えを求めます。\((N-i)^G\) の計算には、繰り返し二乗法(Pythonの pow 関数)を用いることで、各項を \(O(\log G)\) で計算できます。

計算量

  • 時間計算量: \(O(K + M \alpha(K) + N \log G)\)

    • Union-Find の構築に \(O(K + M \alpha(K))\) かかります(\(\alpha\) はアッカーマン関数の逆関数で、実質的に定数とみなせます)。
    • 階乗・逆元の前処理に \(O(N)\) かかります。
    • 包除原理のループが \(N\) 回回り、各ループ内で \(O(\log G)\) の冪乗計算を行うため、この部分は \(O(N \log G)\) です。
    • 全体として、制約 \(N, K, M \leq 2 \times 10^5\) のもとで実行時間制限に十分間に合います。
  • 空間計算量: \(O(K + N)\)

    • Union-Find の配列に \(O(K)\)、階乗および逆元のテーブルに \(O(N)\) のメモリを使用します。

実装のポイント

  • 逆元の高速な前処理(逆元の累積積): 二項係数を繰り返し計算する際、毎回 pow(fact[i], MOD - 2, MOD) を呼び出すと、全体で \(O(N \log \text{MOD})\) かかり、実行時間制限が厳しい言語ではTLE(実行時間制限超過)になる可能性があります。 コードにあるように、まず \(N!\) の逆元を求め、そこから累積的に \((N-1)!, (N-2)!\) の逆元を逆順に求めていく手法を使うことで、すべての逆元を \(O(N)\) で一括して求めることができます。

  • MOD(法)の処理: 包除原理では足し算と引き算が交互に現れます。引き算の際に値が負になる可能性があるため、引き算を行った後は (ans + MOD) % MOD のようにして、常に値が \(0\) 以上 \(MOD-1\) 以下になるように処理します(Pythonでは自動で正の剰余になりますが、他言語での移植時には注意が必要です)。

    ソースコード

import sys

class UnionFind:
    def __init__(self, n):
        self.parent = list(range(n))
        self.count = n

    def find(self, x):
        path = []
        while self.parent[x] != x:
            path.append(x)
            x = self.parent[x]
        for node in path:
            self.parent[node] = x
        return x

    def union(self, x, y):
        rx = self.find(x)
        ry = self.find(y)
        if rx != ry:
            self.parent[ry] = rx
            self.count -= 1
            return True
        return False

def solve():
    input = sys.stdin.read
    data = input().split()
    if not data:
        return
    
    N = int(data[0])
    K = int(data[1])
    M = int(data[2])
    
    uf = UnionFind(K)
    idx = 3
    for _ in range(M):
        u = int(data[idx]) - 1
        v = int(data[idx+1]) - 1
        uf.union(u, v)
        idx += 2
        
    G = uf.count
    
    if G < N:
        print(0)
        return
        
    MOD = 998244353
    
    fact = [1] * (N + 1)
    inv = [1] * (N + 1)
    for i in range(1, N + 1):
        fact[i] = (fact[i-1] * i) % MOD
        
    inv[N] = pow(fact[N], MOD - 2, MOD)
    for i in range(N - 1, -1, -1):
        inv[i] = (inv[i+1] * (i + 1)) % MOD
        
    ans = 0
    for i in range(N + 1):
        val = (fact[N] * inv[i]) % MOD
        val = (val * inv[N-i]) % MOD
        val = (val * pow(N - i, G, MOD)) % MOD
        if i % 2 == 1:
            ans = (ans - val) % MOD
        else:
            ans = (ans + val) % MOD
            
    print(ans)

if __name__ == '__main__':
    solve()

この解説は gemini-3.5-flash-thinking によって生成されました。

posted:
last update: