출시·고도화 중
분할 정복 안내서 · 3/6
이 장에서는 Python으로 두 가지 루틴을 만듭니다. 병합 정렬로 역순 쌍을 세는 함수와 모듈러 빠른 거듭제곱입니다. 이어서 빠른 거듭제곱을 C++, Java, TypeScript로 옮기면서 정수 오버플로를 어떻게 피하는지 살펴봅니다.
아래 코드는 작업용 사본을 임시 버퍼 하나로 제자리 정렬합니다. 단계마다 슬라이싱으로 새 리스트를 만들지 않아도 됩니다.
def count_inversions(a):
"""(i < j 이면서 a[i] > a[j] 인 쌍의 수, a 를 정렬한 사본)을 돌려준다."""
buf = list(a)
tmp = [0] * len(buf)
def solve(lo, hi): # 반열린 구간 [lo, hi)
if hi - lo <= 1:
return 0
mid = (lo + hi) // 2
count = solve(lo, mid) + solve(mid, hi)
i, j, k = lo, mid, lo
while i < mid and j < hi:
if buf[i] <= buf[j]:
tmp[k] = buf[i]
i += 1
else:
tmp[k] = buf[j]
j += 1
count += mid - i # 경계를 넘는 역순 쌍
k += 1
while i < mid:
tmp[k] = buf[i]
i += 1
k += 1
while j < hi:
tmp[k] = buf[j]
j += 1
k += 1
buf[lo:hi] = tmp[lo:hi]
return count
return solve(0, len(buf)), buf
print(count_inversions([5, 2, 4, 7, 1, 3, 2, 6])) # (14, [1, 2, 2, 3, 4, 5, 6, 7])한 줄씩 살펴봅니다.
buf는 사본이므로 호출한 쪽의 리스트는 바뀌지 않습니다. tmp는 모든 병합이 함께 쓰는 임시 버퍼입니다.solve(lo, hi)는 반열린 구간 [lo, hi)를 다룹니다. 길이가 0이나 1인 구간에는 역순 쌍이 없습니다.mid로 구간을 나누고, 두 재귀 호출이 각 절반 안의 역순 쌍을 세면서 두 절반을 정렬해 둡니다.buf[i] <= buf[j]는 값이 같을 때 왼쪽을 먼저 꺼내므로 같은 값을 역순 쌍으로 세지 않습니다.mid - i개 원소 모두보다 작습니다. 이것이 경계를 넘는 역순 쌍입니다.buf[lo:hi] = tmp[lo:hi]로 병합한 결과를 되돌려 씁니다.재귀 깊이는 약 log2 n이므로 원소가 수백만 개여도 Python의 기본 재귀 한도에 걸리지 않습니다.
def power(base, exp, mod):
"""exp >= 0 일 때 base^exp % mod. 재귀 분할 정복."""
if exp == 0:
return 1 % mod
half = power(base, exp // 2, mod)
result = half * half % mod
if exp % 2 == 1:
result = result * base % mod
return result
def power_iter(base, exp, mod):
"""재귀 없이 같은 결과: exp 의 비트를 차례로 읽는다."""
result = 1 % mod
base %= mod
while exp > 0:
if exp & 1:
result = result * base % mod
base = base * base % mod
exp >>= 1
return result
print(power(3, 200, 1_000_000_007), power_iter(3, 200, 1_000_000_007)) # 136318165 136318165exp == 0이 기저 사례입니다. 1 % mod는 mod == 1일 때 0이 되는데, 1로 나눈 나머지는 언제나 0이므로 올바른 답입니다.power(base, exp // 2, mod)가 유일한 재귀 호출입니다. 부분 문제는 정확히 하나입니다.half를 제곱하면 base^(2 * (exp // 2))가 되고, 지수가 홀수면 base를 한 번 더 곱합니다.% mod로 줄여 수가 커지지 않게 합니다.반복 버전은 exp의 비트를 낮은 자리부터 읽습니다. base는 b^1, b^2, b^4, ...로 바뀌고, 현재 비트가 1이면 result에 곱해집니다. 재귀가 없어서 다른 언어로 옮길 때 주로 이 형태를 씁니다. 실제 Python 코드에서는 내장 함수 pow(base, exp, mod)를 쓰면 됩니다.
두 루틴을 믿고 쓰기 전에 무작위 입력 수백 개로 이중 반복문 완전 탐색, 내장 pow와 결과를 비교해 봅니다. 앞 장에서 손으로 추적한 예도 좋은 고정 테스트가 됩니다.
C++에는 기본 큰 정수가 없으므로 두 나머지의 곱이 64비트에 들어가야 합니다. mod가 2^32보다 작으면 result * base는 2^64를 넘지 않습니다.
#include <cstdint>
#include <iostream>
// 1 <= mod < 2^32 이라고 가정한다(모든 곱이 64비트에 들어간다).
std::uint64_t mod_pow(std::uint64_t base, std::uint64_t exp, std::uint64_t mod) {
std::uint64_t result = 1 % mod;
base %= mod;
while (exp > 0) {
if (exp & 1) result = result * base % mod;
base = base * base % mod;
exp >>= 1;
}
return result;
}
int main() {
std::cout << mod_pow(3, 200, 1000000007) << '\n'; // 136318165
}Java의 long은 부호 있는 64비트 정수입니다. mod가 2^31보다 작으면(예: 1_000_000_007) 모든 곱이 2^62 아래에 머뭅니다. Java의 %는 base가 음수면 음수를 돌려줄 수 있으므로 먼저 보정합니다.
public final class ModPow {
// 1 <= mod < 2^31 이라고 가정한다(모든 곱이 long 에 들어간다).
static long modPow(long base, long exp, long mod) {
long result = 1 % mod;
base %= mod;
if (base < 0) base += mod;
while (exp > 0) {
if ((exp & 1) == 1) result = result * base % mod;
base = base * base % mod;
exp >>= 1;
}
return result;
}
public static void main(String[] args) {
System.out.println(modPow(3, 200, 1_000_000_007L)); // 136318165
}
}JavaScript의 number는 배정밀도 실수라서 2^53을 넘으면 정밀도를 잃습니다. 1_000_000_006 * 1_000_000_006만 해도 틀린 값이 나옵니다. bigint를 쓰면 계산이 정확합니다.
export function modPow(base: bigint, exp: bigint, mod: bigint): bigint {
let result = 1n % mod;
base = ((base % mod) + mod) % mod;
while (exp > 0n) {
if (exp & 1n) result = (result * base) % mod;
base = (base * base) % mod;
exp >>= 1n;
}
return result;
}
console.log(modPow(3n, 200n, 1_000_000_007n)); // 136318165n세 언어 모두 Python의 pow(3, 200, 1_000_000_007)과 같은 값을 냅니다.
mid - i를 더합니다.mod를 작게 유지하고, TypeScript에서는 bigint를 씁니다.
댓글 0개
로그인 · 로그인하면 댓글을 남길 수 있습니다.
첫 댓글을 남겨 보세요.