遅延評価:更新をまとめて運ぶ

区間の全要素を直接更新せず、節点に「あとで子へ伝える操作」を保持します。演算を設計してから実装します。

言語:Python / 計算量:適切な代数設計で各操作 O(log N)

前提:セグメント木は区間の集約器

考え方

  1. 区間和なら節点に和と長さを持ちます。
  2. 区間へ x を加えると和は x×長さ 増えます。
  3. 下の節点を読む前に保留した操作を伝播します。

具体例

和10、長さ4の区間に3を加えると新しい和は22。加算操作3と5の合成は8です。

実装

def mapping(add, node):
    total, length = node
    return total + add * length, length

def compose(new, old):
    return new + old

print(mapping(3, (10, 4)))

注意する条件

演算の結合性・作用・合成を確認します。区間代入と加算の混在は順序が重要。講義後半の実装は加算・和に限定。

確認問題

和10、長さ4の区間に3を加えた和は?

解答と理由

22

10 + 3×4 = 22です。

実装課題

操作 f(x)=2x+3 の後に g(x)=4x+5 を適用する合成を求め、x=1 で結果を出す関数を書く。入力なし。

出力:
25
参考実装
def compose(g,f):
    a,b=g;c,d=f
    return a*c,a*d+b
a,b=compose((4,5),(2,3))
print(a*1+b)

区間加算・区間和の完全実装

範囲はすべて0始まりの半開区間 [l,r) に揃えます。data[k]は節点kが担当する区間の和、lazy[k]はその区間の各要素へ加え済みで、子にはまだ伝えていない値です。

区間全体を覆う更新なら、dataには x×区間長、lazyにはxを足して止まれます。一部だけ重なる場合、先にpushで子へ保留分を渡し、必要な子を更新して親を再集計します。

再帰の深さは木の高さ O(log N) です。一般の深いグラフDFSとは違います。更新と取得は O(log N)、メモリ O(N)。この実装は加算と和に限定し、代入やminへそのまま流用しません。

class RangeAddSum:
    def __init__(self, a):
        self.n = len(a)
        self.size = 1
        while self.size < self.n:
            self.size *= 2
        self.data = [0] * (2 * self.size)
        self.lazy = [0] * (2 * self.size)
        self.data[self.size:self.size+self.n] = a
        for k in range(self.size-1, 0, -1):
            self.data[k] = self.data[2*k] + self.data[2*k+1]

    def _apply(self, k, length, x):
        self.data[k] += x * length
        self.lazy[k] += x

    def _push(self, k, length):
        if self.lazy[k] and length > 1:
            x = self.lazy[k]
            self._apply(2*k, length//2, x)
            self._apply(2*k+1, length//2, x)
            self.lazy[k] = 0

    def add(self, l, r, x):
        assert 0 <= l <= r <= self.n
        def visit(k, a, b):
            if r <= a or b <= l:
                return
            if l <= a and b <= r:
                self._apply(k, b-a, x)
                return
            self._push(k, b-a)
            m = (a+b)//2
            visit(2*k, a, m)
            visit(2*k+1, m, b)
            self.data[k] = self.data[2*k] + self.data[2*k+1]
        visit(1, 0, self.size)

    def sum(self, l, r):
        assert 0 <= l <= r <= self.n
        def visit(k, a, b):
            if r <= a or b <= l:
                return 0
            if l <= a and b <= r:
                return self.data[k]
            self._push(k, b-a)
            m = (a+b)//2
            return visit(2*k, a, m) + visit(2*k+1, m, b)
        return visit(1, 0, self.size)

s = RangeAddSum([1, 2, 3, 4])
s.add(1, 4, 5)
print(s.sum(0, 3))  # 16

読了の記録・下書き・メモへ

関連する公式資料・課題