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

C++ 프로그램: 행별 최솟값과 열별 최댓값 수열 쌍의 개수 구하기

문제 소개

세 개의 정수 N, M, K가 주어집니다. N개의 행과 M개의 열로 이루어진 격자의 모든 칸에 1 이상 K 이하의 정수를 하나씩 적는다고 가정해 봅시다. 이때 다음 조건을 만족하는 두 수열 A와 B를 정의합니다.

  • 1부터 N까지의 각 i에 대해, A[i]는 i번째 행에 있는 모든 원소 중 최솟값입니다.
  • 1부터 M까지의 각 j에 대해, B[j]는 j번째 열에 있는 모든 원소 중 최댓값입니다.

구해야 할 것은 가능한 쌍 (A, B)의 개수입니다. 답이 매우 커질 수 있으므로, 결과는 998244353으로 나눈 나머지를 반환합니다.

예를 들어 입력이 N = 2, M = 2, K = 2라면 출력은 7이 됩니다. 실제로 가능한 (A[1], A[2], B[1], B[2])의 조합은 (1,1,1,1), (1,1,1,2), (1,1,2,1), (1,1,2,2), (1,2,2,2), (2,1,2,2), (2,2,2,2)의 7가지뿐입니다.

풀이 전략

이 문제는 다음 단계에 따라 해결할 수 있습니다.

p := 998244353
(a^b) mod p를 반환하는 함수 power(a, b)를 정의한다.
메인 메서드에서 아래를 수행한다.
n이 1이라면:
    power(K, m)을 반환한다.
m이 1이라면:
    power(K, n)을 반환한다.
ans := 0
t := 1부터 t <= K까지 1씩 증가시키며 반복한다:
    ans := (ans + (power(t, n) - power(t - 1, n) + p) mod p * power(K - t + 1, m)) mod p
ans를 반환한다.

핵심 아이디어

이 공식이 성립하는 이유는 다음과 같습니다. 어떤 격자에서든 최댓값을 가지는 행과 최솟값을 가지는 열이 교차하는 칸은 반드시 max(A) 이상이면서 동시에 min(B) 이하이므로, 유효한 격자가 존재하려면 max(A) ≤ min(B)를 만족해야 합니다. 반대로 이 조건을 만족하는 (A, B)에 대해서는 항상 적절한 격자를 구성할 수 있습니다.

따라서 답은 "max(A)가 정확히 t이고 min(B)가 t 이상인 경우"를 t = 1부터 K까지 모두 더한 값과 같습니다. 여기서 max(A) = t가 되는 수열 A의 개수는 tn − (t−1)n이고, min(B) ≥ t가 되는 수열 B의 개수는 (K−t+1)m입니다. 두 값을 곱한 뒤 모두 더하면 원하는 답을 얻을 수 있습니다.

한편 n = 1이면 A의 유일한 원소는 전체 격자의 최솟값, 즉 min(B)로 자동으로 고정되므로 서로 다른 B의 개수 Km이 곧 답이 됩니다. m = 1인 경우도 같은 논리로 Kn입니다.

C++ 구현 예제

아래 구현을 통해 더 잘 이해해 봅시다.

#include <bits/stdc++.h>
using namespace std;

long p = 998244353;

long power(long a, long b, long ret = 1){
    for (; b; b >>= 1, a = a * a % p)
       if (b & 1)
          ret = ret * a % p;
    return ret;
}
long solve(int n, int m, int K){
    if (n == 1)
       return power(K, m);
    if (m == 1)
       return power(K, n);
    long ans = 0;
    for (long t = 1; t <= K; t++){
       ans = (ans + (power(t, n) - power(t - 1, n) + p) % p * power(K - t + 1, m)) % p;
    }
    return ans;
}
int main(){
   int N = 2;
   int M = 2;
   int K = 2;
   cout << solve(N, M, K) << endl;
}

power() 함수는 거듭제곱 분할 정복(빠른 거듭제곱) 기법을 사용해 O(log b) 시간 안에 (ab) mod p를 계산합니다. 따라서 전체 시간 복잡도는 O(K log N) 수준으로 효율적입니다.

입력

2, 2, 2

출력

7