Computer >> 컴퓨터 >  >> 프로그래밍 >> C++

두 가지 조건을 만족하는 블록 색칠 방법의 수를 세는 C++ 프로그램

문제 개요

세 개의 정수 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