import heapq

class SparseTable:

    def __init__(self, arr, mode="max"):
        assert mode in ("max", "min")
        self.n = len(arr)
        self.mode = mode
        self.op = max if mode == "max" else min

        self.log = [0] * (self.n + 1)
        for i in range(2, self.n + 1):
            self.log[i] = self.log[i // 2] + 1

        k = self.log[self.n] + 1 if self.n > 0 else 1
        self.table = [[0] * self.n for _ in range(k)]
        self.table[0] = list(arr)

        for j in range(1, k):
            half = 1 << (j - 1)
            for i in range(self.n - (1 << j) + 1):
                self.table[j][i] = self.op(
                    self.table[j - 1][i],
                    self.table[j - 1][i + half],
                )

    def query(self, l, r):
        assert 0 <= l <= r < self.n
        j = self.log[r - l + 1]
        return self.op(
            self.table[j][l],
            self.table[j][r - (1 << j) + 1],
        )

class Solution:
    def maxTotalValue(self, nums: List[int], k: int) -> int:
        n = len(nums)
        max_query = SparseTable(nums, "max")
        min_query = SparseTable(nums, "min")

        pq = []
        seen = set()
        query = lambda l, r: seen.add((l, r)) or (min_query.query(l, r) - max_query.query(l, r), l, r)

        for i in range(n):
            pq.append(query(i, n - 1))

        def _next():  
            if pq:
                v, l, r = heapq.heappop(pq)
                if r - l > 1 and (l + 1, r) not in seen:
                    heapq.heappush(pq, query(l + 1, r))
                if r - l > 1 and (l, r - 1) not in seen:
                    heapq.heappush(pq, query(l, r - 1))
                return -v
            return 0
        
        return sum(_next() for i in range(k))
