Computer >> 컴퓨터 >  >> 프로그래밍 >> Python

Python으로 표현식 결과의 최대 빈도 기댓값 구하는 프로그램


M개의 서로 다른 표현식이 있고, 각 표현식의 답은 1부터 N까지(양 끝값 포함) 범위 안에 있다고 가정해 보겠습니다. 이때 1부터 N까지 각 숫자 i의 발생 빈도를 f(i)라 하고, x = max(f(i))로 정의하면, 우리가 구해야 하는 것은 바로 이 x의 기댓값(expected value)입니다.

예를 들어 입력이 M = 3, N = 3이라면 출력은 2.2가 됩니다. 가능한 모든 시퀀스와 각 시퀀스에서 가장 자주 등장한 숫자의 빈도(최대 빈도)는 아래 표와 같습니다.

시퀀스최대 빈도
1113
1122
1132
1222
1231
1331
2223
2232
2332
3333

전체 10가지 경우 중 최대 빈도가 3인 경우는 111, 222, 333의 3가지, 2인 경우는 6가지, 1인 경우는 123 한 가지뿐입니다. 따라서 기댓값은 다음과 같이 계산할 수 있습니다.

$$E(x) = \sum P(x) * x = P(1) + 2P(2) + 3P(3) = \frac{1}{10} + 2 * \frac{6}{10} + 3 * \frac{3}{10} = \frac{22}{10}$$

문제 해결 접근 방법

이 문제는 조합(combination) 계산과 포함-배제 원리(inclusion-exclusion principle)를 활용하면 효율적으로 해결할 수 있습니다. 먼저 중복 계산을 피하기 위해 메모이제이션을 적용한 nCr() 함수를 준비합니다.

  • combination := 계산 결과를 저장할 새로운 맵(딕셔너리)
  • nCr(n, k_in) 함수를 정의합니다.
  • k := k_in과 (n - k_in) 중 더 작은 값
  • n < k 또는 k < 0이면 0을 반환합니다.
  • (n, k)가 combination에 이미 존재하면 저장된 값을 반환합니다.
  • k == 0이면 1을 반환합니다.
  • n == k이면 1을 반환합니다.
  • 그 외의 경우에는 a := 1로 초기화한 뒤, cnt를 0부터 k-1까지 반복하면서 a에 (n - cnt)를 곱하고 (cnt + 1)로 나눈 몫을 저장하며, 각 단계의 결과를 combination[(n, cnt + 1)]에 기록한 후 마지막에 a를 반환합니다.

메인 로직은 다음 단계로 진행됩니다.

  • arr := 새로운 리스트
  • k를 2부터 M + 1까지 반복합니다.
    • a := 1, s := 0으로 초기화합니다.
    • i를 0부터 M // k + 2까지 반복합니다.
      • M < i * k이면 반복을 종료합니다.
      • s := s + a * nCr(N, i) * nCr(N - 1 + M - i * k, M - i * k)
      • a := -a (부호를 번갈아 변경하여 포함-배제 원리를 적용)
    • s를 arr의 끝에 추가합니다.
  • total := arr의 마지막 요소(전체 경우의 수)
  • diff := arr[0]으로 시작하고, 이후에는 인접 요소의 차(arr[cnt + 1] - arr[cnt])로 구성된 차분 배열
  • output := sum(diff[cnt] * (cnt + 1) / total)
  • output을 반환합니다.

예시 코드

더 나은 이해를 위해 다음 구현을 살펴보겠습니다.

combination = {}
def nCr(n, k_in):
    k = min(k_in, n - k_in)
    if n < k or k < 0:
        return 0
    elif (n, k) in combination:
        return combination[(n, k)]
    elif k == 0:
        return 1
    elif n == k:
        return 1
    else:
        a = 1
        for cnt in range(k):
            a *= (n - cnt)
            a //= (cnt + 1)
            combination[(n, cnt + 1)] = a
        return a

def solve(M, N):
    arr = []
    for k in range(2, M + 2):
        a = 1
        s = 0
        for i in range(M // k + 2):
            if (M < i * k):
                break
            s += a * nCr(N, i) * nCr(N - 1 + M - i * k, M - i * k)
            a *= -1
        arr.append(s)
    total = arr[-1]
    diff = [arr[0]] + [arr[cnt + 1] - arr[cnt] for cnt in range(M - 1)]
    output = sum(diff[cnt] * (cnt + 1) / total for cnt in range(M))
    return output

M = 3
N = 3
print(solve(M, N))

입력

3, 3

출력

2.2