문제 개요
세 개의 정수 N, M, K가 주어진 상황을 생각해 봅시다. N개의 블록이 한 줄로 나열되어 있으며, 다음 두 가지 조건에 따라 블록을 칠하는 방법의 수를 구하려고 합니다. 두 가지 칠하기 결과는 블록들의 색상 배치가 서로 다를 때에만 다른 것으로 간주됩니다.
- 각 블록에는 M가지 색상 중 하나를 골라 칠합니다. (모든 색상을 반드시 사용할 필요는 없습니다)
- 같은 색으로 칠해진 인접한 블록 쌍은 최대 K쌍까지만 존재할 수 있습니다.
답이 너무 커질 수 있으므로, 결과는 998244353으로 나눈 나머지를 반환합니다.
예를 들어 입력이 N = 3, M = 2, K = 1이라면 출력은 6입니다. 112, 121, 122, 211, 212, 221의 여섯 가지 방식으로 블록을 칠할 수 있기 때문입니다.
접근 방법
이 문제는 조합론적 접근과 모듈러 거듭제곱을 활용하면 효율적으로 해결할 수 있습니다. 핵심 아이디어는 다음과 같습니다.
N개의 블록 사이에는 총 N−1개의 인접 경계가 있습니다. 이 경계 중 정확히 i개를 '양옆 블록이 같은 색'인 지점으로 선택하면, 블록 열은 N−i개의 연속된 색상 구간으로 나뉩니다. 첫 번째 구간은 M가지 색 중 자유롭게 고를 수 있고, 이후 각 구간은 바로 앞 구간과 다른 색이어야 하므로 M−1가지의 선택지가 있습니다. 따라서 인접한 같은 색 쌍이 정확히 i개인 경우의 수는 다음과 같습니다.
C(N−1, i) × M × (M−1)^(N−i−1)
i를 0부터 K까지 모두 더하면 '최대 K쌍'이라는 조건을 만족하는 전체 경우의 수를 얻을 수 있습니다.
큰 수의 이항계수를 빠르게 계산하기 위해 팩토리얼 배열 fac와 그 모듈러 역원 배열 inv를 미리 구해 둡니다. 역원은 페르마의 소정리를 이용해 거듭제곱 형태로 계산합니다. 전체 알고리즘 단계는 다음과 같습니다.
maxm := 2×10^6 + 5
p := 998244353
크기가 maxm인 두 배열 fac와 inv를 선언합니다.
ppow() 함수를 정의합니다. 매개변수는 a, b, p입니다.
ans := 1 mod p
a := a mod p
b가 0이 아닌 동안 다음을 반복합니다:
b가 홀수이면:
ans := ans * a mod p
a := a * a mod p
b := b / 2
ans를 반환합니다.
C() 함수를 정의합니다. 매개변수는 n, m입니다.
m < 0 또는 m > n이면:
0을 반환합니다.
fac[n] * inv[m] mod p * inv[n - m] mod p를 반환합니다.
메인 메서드에서 다음을 수행합니다.
fac[0] := 1
i := 1부터 i < maxm까지 (i는 1씩 증가):
fac[i] := fac[i - 1] * i mod p
inv[maxm - 1] := ppow(fac[maxm - 1], p - 2, p)
i := maxm - 2부터 i >= 0까지 (i는 1씩 감소):
inv[i] := (i + 1) * inv[i + 1] mod p
ans := 0
i := 0부터 i <= k까지 (i는 1씩 증가):
t := C(n - 1, i)
tt := m * ppow(m - 1, n - i - 1, p)
ans := (ans + t * tt mod p) mod p
ans를 반환합니다.
예제 코드 (C++)
아래 구현을 통해 더 잘 이해해 볼 수 있습니다.
#include <bits/stdc++.h>
using namespace std;
const long maxm = 2e6 + 5;
const long p = 998244353;
long fac[maxm], inv[maxm];
long ppow(long a, long b, long p){
long ans = 1 % p;
a %= p;
while (b){
if (b & 1)
ans = ans * a % p;
a = a * a % p;
b >>= 1;
}
return ans;
}
long C(long n, long m){
if (m < 0 || m > n)
return 0;
return fac[n] * inv[m] % p * inv[n - m] % p;
}
long solve(long n, long m, long k){
fac[0] = 1;
for (long i = 1; i < maxm; i++)
fac[i] = fac[i - 1] * i % p;
inv[maxm - 1] = ppow(fac[maxm - 1], p - 2, p);
for (long i = maxm - 2; i >= 0; i--)
inv[i] = (i + 1) * inv[i + 1] % p;
long ans = 0;
for (long i = 0; i <= k; i++){
long t = C(n - 1, i);
long tt = m * ppow(m - 1, n - i - 1, p) % p;
ans = (ans + t * tt % p) % p;
}
return ans;
}
int main(){
int N = 3;
int M = 2;
int K = 1;
cout << solve(N, M, K) << endl;
}
입력
3, 2, 1
출력
6