已發布·持續改進
線段樹 指南 · 5/6
本章目前僅提供英文版。
These problems were written for this guide to cover the classic segment tree and Fenwick tree patterns. Try each one before reading the approach. All solutions use plain Python.
A store's sales for n days are given. ("set", d, v) corrects day d to v, and ("sum", l, r) reports total sales for days l through r-1. Both n and the request count are at most 200,000.
Approach: range sums make a Fenwick tree the simplest fit. It only supports additions, so keep the original array and add the difference. Building with n calls to add costs O(n log n); pushing each cell into its parent, as below, takes O(n).
def solve(sales, requests):
n = len(sales)
tree = [0] + sales[:]
for i in range(1, n + 1): # O(n) build
parent = i + (i & -i)
if parent <= n:
tree[parent] += tree[i]
def add(i, delta):
i += 1
while i <= n:
tree[i] += delta
i += i & -i
def prefix(count):
total = 0
while count > 0:
total += tree[count]
count -= count & -count
return total
out = []
for kind, x, y in requests:
if kind == "set":
add(x, y - sales[x])
sales[x] = y
else:
out.append(prefix(y) - prefix(x))
return out
print(solve([4, 0, 7, 2, 5], [("sum", 1, 4), ("set", 1, 6), ("sum", 0, 5)])) # [9, 24]n sensors report temperatures. ("set", i, t) changes sensor i to t, and ("min", l, r) asks for the lowest reading among sensors l through r-1.
Approach: a minimum cannot be undone by subtraction, so use a segment tree with min as the operation and infinity as the identity.
def coldest(temps, requests):
n = len(temps)
INF = float("inf")
tree = [INF] * n + temps
for i in range(n - 1, 0, -1):
tree[i] = min(tree[2 * i], tree[2 * i + 1])
out = []
for kind, x, y in requests:
if kind == "set":
i = x + n
tree[i] = y
while i > 1:
i //= 2
tree[i] = min(tree[2 * i], tree[2 * i + 1])
else:
lo, hi, best = x + n, y + n, INF
while lo < hi:
if lo & 1:
best = min(best, tree[lo]); lo += 1
if hi & 1:
hi -= 1; best = min(best, tree[hi])
lo //= 2; hi //= 2
out.append(best)
return out
print(coldest([3, -2, 5, 0, 1], [("min", 2, 5), ("set", 3, 4), ("min", 2, 5), ("min", 0, 5)]))
# [0, 1, -2]Count the pairs i < j with a[i] > a[j] (inversions) in an array of up to 200,000 integers with a very wide value range.
Approach: scan left to right, adding how many earlier values are larger. Replace values with their sorted ranks (coordinate compression) and keep a count per rank in a Fenwick tree: O(n log n) in total.
def count_inversions(a):
ranks = {v: i + 1 for i, v in enumerate(sorted(set(a)))}
m = len(ranks)
tree = [0] * (m + 1)
seen = inversions = 0
for v in a:
r = ranks[v]
smaller_or_equal, i = 0, r
while i > 0:
smaller_or_equal += tree[i]
i -= i & -i
inversions += seen - smaller_or_equal
while r <= m:
tree[r] += 1
r += r & -r
seen += 1
return inversions
print(count_inversions([8, 4, 2, 1])) # 6
print(count_inversions([3, 1, 2, 3, 1])) # 5("add", l, r, v) adds v to the prices of products l through r-1 (negative for a discount), and ("sum", l, r) reports their total.
Approach: range updates plus range queries call for lazy propagation. At a fully covered node, add v * length to its sum, accumulate v in its lazy slot, and stop. push hands the lazy value to the children only when you must descend.
def solve_lazy(prices, requests):
n = len(prices)
tree, lazy = [0] * (4 * n), [0] * (4 * n)
def build(node, lo, hi):
if hi - lo == 1:
tree[node] = prices[lo]
return
mid = (lo + hi) // 2
build(2 * node, lo, mid); build(2 * node + 1, mid, hi)
tree[node] = tree[2 * node] + tree[2 * node + 1]
def apply(node, lo, hi, v):
tree[node] += v * (hi - lo)
lazy[node] += v
def push(node, lo, hi):
if lazy[node]:
mid = (lo + hi) // 2
apply(2 * node, lo, mid, lazy[node])
apply(2 * node + 1, mid, hi, lazy[node])
lazy[node] = 0
def add(node, lo, hi, l, r, v):
if r <= lo or hi <= l:
return
if l <= lo and hi <= r:
apply(node, lo, hi, v)
return
push(node, lo, hi)
mid = (lo + hi) // 2
add(2 * node, lo, mid, l, r, v); add(2 * node + 1, mid, hi, l, r, v)
tree[node] = tree[2 * node] + tree[2 * node + 1]
def total(node, lo, hi, l, r):
if r <= lo or hi <= l:
return 0
if l <= lo and hi <= r:
return tree[node]
push(node, lo, hi)
mid = (lo + hi) // 2
return total(2 * node, lo, mid, l, r) + total(2 * node + 1, mid, hi, l, r)
build(1, 0, n)
out = []
for req in requests:
if req[0] == "add":
add(1, 0, n, req[1], req[2], req[3])
else:
out.append(total(1, 0, n, req[1], req[2]))
return out
print(solve_lazy([10, 20, 30, 40, 50], [("add", 1, 4, -5), ("sum", 0, 3), ("sum", 2, 5)]))
# [50, 110]Range sums with point updates suggest a Fenwick tree, minimums or maximums a segment tree, and range updates a lazy segment tree. Counting questions such as "how many earlier values are larger" often combine coordinate compression with a Fenwick tree.
0 則留言
登入 · 登入後即可留言。
來留下第一則留言吧。