range-sum-query-mutable-v2.py (2516B)
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 update_val(root, idx, delta): 35 if root is None: 36 return 37 if idx < root.rng[0] or idx > root.rng[1]: 38 return 39 40 root.rv += delta 41 update_val(root.right, idx, delta) 42 update_val(root.left, idx, delta) 43 44 45 46 47 def make_segment(nums): 48 forest = [] 49 50 idx = 0 51 for num in nums: 52 current = Node() 53 current.rng = [idx,idx] 54 current.rv = num 55 forest.append(current) 56 idx += 1 57 58 while len(forest) > 1: 59 nf = [] 60 61 for i in range(0, len(forest), 2): 62 if i + 1 < len(forest): 63 nf.append(combine(forest[i], forest[i + 1])) 64 else: 65 nf.append(forest[i]) 66 67 forest = nf 68 69 return forest[0] 70 71 class NumArray: 72 73 # len(nums) is at most 30_000 74 def __init__(self, nums: List[int]): 75 self.nums = nums 76 self.segment_tree = make_segment(nums) 77 78 # nums[index] = val 79 def update(self, index: int, val: int) -> None: 80 # could do away with this list if we accept log(n) cost for lookups here 81 # this would retain the same WCTC of log(n) for updating but increase the constant. 82 # this would decrease memory usage, but still in O(n) for that too. 83 prior = self.nums[index] 84 self.nums[index] = val 85 update_val(self.segment_tree,index,val - prior) 86 return 87 88 # range queries can be size of nums 89 # precondition: left <= right 90 # return: sum(nums[left,right]) - inclusive of left and right 91 def sumRange(self, left: int, right: int) -> int: 92 return find_sum([left,right], self.segment_tree) 93 94 # Your NumArray object will be instantiated and called as such: 95 # obj = NumArray(nums) 96 # obj.update(index,val) 97 # param_2 = obj.sumRange(left,right)