출시·고도화 중
탐욕 알고리즘 안내서 · 3/6
이 장에서는 탐욕 알고리즘의 대표 두 가지를 Python으로 깔끔하게 구현합니다. 하나는 겹치지 않는 구간을 가장 많이 고르는 구간 스케줄링이고, 다른 하나는 기호마다 허프만 부호의 길이를 구하는 함수입니다. 이어서 구간 스케줄링을 C++, Java, TypeScript로 옮겨 언어마다 정렬과 비교를 어떻게 쓰는지 비교합니다.
import heapq
def max_non_overlapping(intervals: list[tuple[int, int]]) -> list[tuple[int, int]]:
"""서로 겹치지 않는 반열린 구간 [start, end)를 가장 많이 고른다."""
chosen: list[tuple[int, int]] = []
last_end = float("-inf")
for start, end in sorted(intervals, key=lambda iv: iv[1]):
if start >= last_end:
chosen.append((start, end))
last_end = end
return chosen
def huffman_code_lengths(freq: dict[str, int]) -> dict[str, int]:
"""기호마다 허프만 부호의 비트 길이를 돌려준다."""
symbols = list(freq)
if len(symbols) <= 1:
return {s: 1 for s in symbols}
heap = [(w, i) for i, w in enumerate(freq.values())]
heapq.heapify(heap)
parent = [-1] * len(symbols) # 노드 번호 -> 부모 번호, n 이상은 내부 노드
while len(heap) > 1:
w1, a = heapq.heappop(heap)
w2, b = heapq.heappop(heap)
node = len(parent)
parent.append(-1)
parent[a] = parent[b] = node
heapq.heappush(heap, (w1 + w2, node))
depth = [0] * len(parent)
for node in range(len(parent) - 2, -1, -1):
depth[node] = depth[parent[node]] + 1
return {s: depth[i] for i, s in enumerate(symbols)}sorted(intervals, key=lambda iv: iv[1]): 끝 시간 기준으로 정렬한 새 리스트를 만듭니다. 입력을 바꾸지 않으므로 호출한 쪽의 데이터가 안전합니다.last_end = float("-inf"): 아직 고른 구간이 없으므로 어떤 구간도 받아들일 수 있게 음의 무한대로 시작합니다.if start >= last_end: 반열린 구간이라 직전 구간이 끝나는 시각에 바로 시작해도 겹치지 않습니다. 닫힌 구간이라면 >로 바꿉니다.last_end에 기록하면, 다음 후보는 이 값과만 비교하면 됩니다. 정렬 덕분에 이미 고른 다른 구간과는 비교할 필요가 없습니다.(무게, 노드 번호)를 넣습니다. 무게가 같을 때 번호로 비교되므로 결과가 항상 같고, 비교할 수 없는 객체를 넣다가 TypeError가 나는 일도 없습니다.a, b를 꺼내 새 내부 노드 node의 자식으로 만들고, 합친 무게로 다시 힙에 넣습니다. 이것이 탐욕 선택입니다.from collections import Counter
meetings = [(1, 3), (2, 5), (4, 7), (1, 8), (6, 9), (8, 10)]
print(max_non_overlapping(meetings)) # [(1, 3), (4, 7), (8, 10)]
lengths = huffman_code_lengths({"a": 5, "b": 9, "c": 12, "d": 13, "e": 16, "f": 45})
print(lengths) # {'a': 4, 'b': 4, 'c': 3, 'd': 3, 'e': 3, 'f': 1}
text = "abracadabra"
counts = Counter(text)
bits = sum(counts[s] * n for s, n in huffman_code_lengths(counts).items())
print(bits, "비트, 고정 8비트라면", 8 * len(text)) # 23 비트, 고정 8비트라면 88#include <algorithm>
#include <iostream>
#include <limits>
#include <utility>
#include <vector>
using Interval = std::pair<long long, long long>; // [start, end)
std::vector<Interval> maxNonOverlapping(std::vector<Interval> intervals) {
std::sort(intervals.begin(), intervals.end(),
[](const Interval& a, const Interval& b) { return a.second < b.second; });
std::vector<Interval> chosen;
long long lastEnd = std::numeric_limits<long long>::min();
for (const auto& [start, end] : intervals) {
if (start >= lastEnd) {
chosen.push_back({start, end});
lastEnd = end;
}
}
return chosen;
}
int main() {
std::vector<Interval> meetings{{1, 3}, {2, 5}, {4, 7}, {1, 8}, {6, 9}, {8, 10}};
for (const auto& [s, e] : maxNonOverlapping(meetings)) std::cout << s << ' ' << e << '\n';
}벡터를 값으로 받아 복사본을 정렬하므로 호출한 쪽의 벡터는 그대로입니다. 구조적 바인딩(auto& [start, end])은 C++17부터 쓸 수 있습니다.
import java.util.ArrayList;
import java.util.Arrays;
import java.util.Comparator;
import java.util.List;
public class IntervalScheduling {
record Interval(long start, long end) {}
static List<Interval> maxNonOverlapping(Interval[] intervals) {
Interval[] sorted = intervals.clone();
Arrays.sort(sorted, Comparator.comparingLong(Interval::end));
List<Interval> chosen = new ArrayList<>();
long lastEnd = Long.MIN_VALUE;
for (Interval iv : sorted) {
if (iv.start() >= lastEnd) {
chosen.add(iv);
lastEnd = iv.end();
}
}
return chosen;
}
public static void main(String[] args) {
Interval[] meetings = {
new Interval(1, 3), new Interval(2, 5), new Interval(4, 7),
new Interval(1, 8), new Interval(6, 9), new Interval(8, 10),
};
System.out.println(maxNonOverlapping(meetings));
}
}record는 Java 16부터 쓸 수 있습니다. Comparator.comparingLong은 뺄셈으로 비교할 때 생기는 오버플로 걱정이 없습니다.
type Interval = readonly [start: number, end: number];
export function maxNonOverlapping(intervals: readonly Interval[]): Interval[] {
const sorted = [...intervals].sort((a, b) => a[1] - b[1]);
const chosen: Interval[] = [];
let lastEnd = -Infinity;
for (const [start, end] of sorted) {
if (start >= lastEnd) {
chosen.push([start, end]);
lastEnd = end;
}
}
return chosen;
}
console.log(maxNonOverlapping([[1, 3], [2, 5], [4, 7], [1, 8], [6, 9], [8, 10]]));Array.prototype.sort는 원본을 바꾸므로 펼침 연산자로 복사한 뒤 정렬합니다. 시간 값이 안전한 정수 범위 안에 있으면 a[1] - b[1] 비교로 충분합니다.
>=), 원본을 바꾸지 않는 복사만 지키면 됩니다.
댓글 0개
로그인 · 로그인하면 댓글을 남길 수 있습니다.
첫 댓글을 남겨 보세요.