algorithms

Algorithm implementations
git clone git://git.laack.co/algorithms.git
Log | Files | Refs | README

range-sum-query-mutable.py (1991B)


      1 class Node:
      2     def __init__(self):
      3         self.rng = [0, 0]
      4         self.rv = 0
      5         self.left = None
      6         self.right = None
      7 
      8 def combine(a,b):
      9     root = Node()
     10     root.left = a
     11     root.right = b
     12     root.rv = a.rv + b.rv
     13     root.rng[0] = min(a.rng[0], b.rng[0])
     14     root.rng[1] = max(a.rng[1], b.rng[1])
     15     return root
     16 
     17 
     18 # cases:
     19     # if no overlap between rng and root return
     20     # if same, return val
     21     # if different return sum of left + right
     22 
     23 def find_sum(rng, root):
     24     if root is None:
     25         return 0
     26     if rng[1] < root.rng[0] or rng[0] > root.rng[1]:
     27         return 0
     28     
     29     if root.rng[0] >= rng[0] and root.rng[1] <= rng[1]:
     30         return root.rv
     31     
     32     return find_sum(rng,root.left) + find_sum(rng,root.right)
     33 
     34 def make_segment(nums):
     35     forest = []
     36 
     37     idx = 0
     38     for num in nums:
     39         current = Node()
     40         current.rng = [idx,idx]
     41         current.rv = num
     42         forest.append(current)
     43         idx += 1
     44     
     45     while len(forest) > 1:
     46         nf = []
     47 
     48         for i in range(0, len(forest), 2):
     49             if i + 1 < len(forest):
     50                 nf.append(combine(forest[i], forest[i + 1]))
     51             else:
     52                 nf.append(forest[i])
     53 
     54         forest = nf
     55 
     56     return forest[0]
     57 
     58 class NumArray:
     59 
     60     # len(nums) is at most 30_000
     61     def __init__(self, nums: List[int]):
     62         self.nums = nums
     63         self.segment_tree = make_segment(nums)
     64         
     65     # nums[index] = val 
     66     def update(self, index: int, val: int) -> None:
     67         self.nums[index] = val
     68         self.segment_tree = make_segment(self.nums)
     69         return
     70 
     71     # range queries can be size of nums
     72     # precondition: left <= right
     73     # return: sum(nums[left,right]) - inclusive of left and right
     74     def sumRange(self, left: int, right: int) -> int:
     75         return find_sum([left,right], self.segment_tree)
     76 
     77 # Your NumArray object will be instantiated and called as such:
     78 # obj = NumArray(nums)
     79 # obj.update(index,val)
     80 # param_2 = obj.sumRange(left,right)