Use this to learn the idea, then write your own version.
123 4import heapq5 6 78class Solution(object):9 def maxTotalValue(self, nums, k):10 """11 :type nums: List[int]12 :type k: int13 :rtype: int14 """15 def nxt(left, right, i, j):16 while not (left <= idxs[i] <= right):17 i += 118 while not (left <= idxs[j] <= right):19 j -= 120 return (i, j)21 22 idxs = range(len(nums))23 idxs.sort(key=lambda x: (nums[x], x))24 lookup = {(0, len(nums)-1):(0, len(idxs)-1)}25 max_heap = [(-(nums[idxs[len(idxs)-1]]-nums[idxs[0]]), (0, len(idxs)-1))]26 result = 027 while k:28 v, (l, r) = heapq.heappop(max_heap)29 i, j = lookup[(l, r)]30 nl, nr = min(idxs[i], idxs[j]), max(idxs[i], idxs[j])31 c = min((nl-l+1)*(r-nr+1), k)32 k -= c33 result += c*(-v)34 if nl+1 <= r and (nl+1, r) not in lookup:35 lookup[(nl+1, r)] = (ni, nj) = nxt(nl+1, r, i, j)36 heapq.heappush(max_heap, (-(nums[idxs[nj]]-nums[idxs[ni]]), (nl+1, r)))37 if l <= nr-1 and (l, nr-1) not in lookup:38 lookup[(l, nr-1)] = (ni, nj) = nxt(l, nr-1, i, j)39 heapq.heappush(max_heap, (-(nums[idxs[nj]]-nums[idxs[ni]]), (l, nr-1)))40 return result41 42 434445import heapq46 47 4849class Solution2(object):50 def maxTotalValue(self, nums, k):51 """52 :type nums: List[int]53 :type k: int54 :rtype: int55 """56 57 58 59 60 61 class SparseTable(object):62 def __init__(self, arr, fn):63 self.fn = fn64 self.bit_length = [0]65 n = len(arr)66 k = n.bit_length()-1 67 for i in xrange(k+1):68 self.bit_length.extend(i+1 for _ in xrange(min(1<<i, (n+1)-len(self.bit_length))))69 self.st = [[0]*n for _ in xrange(k+1)]70 self.st[0] = arr[:]71 for i in xrange(1, k+1): 72 for j in xrange((n-(1<<i))+1):73 self.st[i][j] = fn(self.st[i-1][j], self.st[i-1][j+(1<<(i-1))])74 75 def query(self, L, R): 76 i = self.bit_length[R-L+1]-1 77 return self.fn(self.st[i][L], self.st[i][R-(1<<i)+1])78 79 rmq_min = SparseTable(nums, min)80 rmq_max = SparseTable(nums, max)81 max_heap = [(-(rmq_max.query(i, len(nums)-1)-rmq_min.query(i, len(nums)-1)), (i, len(nums)-1)) for i in xrange(len(nums))]82 heapq.heapify(max_heap)83 result = 084 for _ in xrange(k):85 v, (i, j) = heappop(max_heap)86 result += -v87 if i <= j-1:88 heapq.heappush(max_heap, (-(rmq_max.query(i, j-1)-rmq_min.query(i, j-1)), (i, j-1)))89 return result90 91 929394import heapq95 96 9798class Solution3(object):99 def maxTotalValue(self, nums, k):100 """101 :type nums: List[int]102 :type k: int103 :rtype: int104 """105 class SegmentTree(object):106 def __init__(self, N, build_fn, query_fn):107 self.tree = [None]*(1<<((N-1).bit_length()+1))108 self.base = len(self.tree)>>1109 self.query_fn = query_fn110 for i in xrange(self.base, self.base+N):111 self.tree[i] = build_fn(i-self.base)112 for i in reversed(xrange(1, self.base)):113 self.tree[i] = query_fn(self.tree[i<<1], self.tree[(i<<1)+1])114 115 def query(self, L, R):116 if L > R:117 return None118 L += self.base119 R += self.base120 left = right = None121 while L <= R:122 if L & 1:123 left = self.query_fn(left, self.tree[L])124 L += 1125 if R & 1 == 0:126 right = self.query_fn(self.tree[R], right)127 R -= 1128 L >>= 1129 R >>= 1130 return self.query_fn(left, right)131 132 133 st_min = SegmentTree(len(nums), build_fn=lambda x: nums[x], query_fn=lambda x, y: y if x is None else x if y is None else min(x, y))134 st_max = SegmentTree(len(nums), build_fn=lambda x: nums[x], query_fn=lambda x, y: y if x is None else x if y is None else max(x, y))135 max_heap = [(-(st_max.query(i, len(nums)-1)-st_min.query(i, len(nums)-1)), (i, len(nums)-1)) for i in xrange(len(nums))]136 heapq.heapify(max_heap)137 result = 0138 for _ in xrange(k):139 v, (i, j) = heappop(max_heap)140 result += -v141 if i <= j-1:142 heapq.heappush(max_heap, (-(st_max.query(i, j-1)-st_min.query(i, j-1)), (i, j-1)))143 return result144