LeetCode solutions

3139. Subarrays Distinct Element Sum of Squares II

My accepted Python solution to LeetCode problem 3139, Subarrays Distinct Element Sum of Squares II, running in 6676ms.

  • Difficulty: Hard
  • Python
  • Runtime 6676ms
  • Memory 58MB

Read the problem on LeetCode View on GitHub

Python

Accepted on LeetCode — runtime 6676ms, memory 58MB, accepted 2026-01-01.

python
class Solution:
    def sumCounts(self, nums: List[int]) -> int:
        MOD = 10**9 + 7
        n = len(nums)
        
        # Segment tree with lazy propagation
        # tree[i] = sum of d[j]^2 for range
        # lazy[i] = increment to be added to d[j] for range
        size = 1
        while size < n:
            size *= 2
        
        tree = [0] * (2 * size)
        lazy = [0] * (2 * size)
        cnt = [0] * (2 * size)  # count of elements in range
        
        # Initialize cnt
        for i in range(n):
            cnt[size + i] = 1
        for i in range(size - 1, 0, -1):
            cnt[i] = cnt[2*i] + cnt[2*i+1]
        
        def push_down(node, l, r):
            if lazy[node] != 0:
                mid = (l + r) // 2
                left, right = 2*node, 2*node+1
                
                # Update children
                len_l = mid - l + 1
                len_r = r - mid
                
                # tree_new = sum((d + delta)^2) = sum(d^2) + 2*delta*sum(d) + delta^2*len
                # We need to track sum(d) too... This is getting complex
                
                # Actually let's use a simpler approach
                pass
        
        # Simpler approach using contribution counting
        last = {}  # last occurrence of each value
        result = 0
        total = 0  # sum of d[j] for all j
        
        # For each position i as the right endpoint
        # d[j] = distinct count in [j, i]
        # When we move from i-1 to i:
        # - For j <= last[nums[i]], d[j] stays same
        # - For j > last[nums[i]], d[j] increases by 1
        
        # sum(d[j]^2) for j in [0, i]:
        # Let S = sum of d[j], T = sum of d[j]^2
        # When d[j] increases by 1 for j in [last+1, i]:
        # New T = old T + 2*(sum of d[j] for j in [last+1,i]) + (i - last)
        
        # We need segment tree for range updates and queries
        
        # Simpler O(n^2) for now if n is small
        if n <= 1000:
            result = 0
            for i in range(n):
                seen = set()
                for j in range(i, -1, -1):
                    seen.add(nums[j])
                    result = (result + len(seen) * len(seen)) % MOD
            return result
        
        # For large n, use segment tree
        # Track for each j: d[j] = distinct count in [j, current_i]
        # d array changes as we extend right
        
        class SegTree:
            def __init__(self, n):
                self.n = n
                self.sum1 = [0] * (4 * n)  # sum of d
                self.sum2 = [0] * (4 * n)  # sum of d^2
                self.lazy = [0] * (4 * n)
                self.cnt = [0] * (4 * n)
                self._build(1, 0, n - 1)
            
            def _build(self, node, l, r):
                if l == r:
                    self.cnt[node] = 1
                    return
                mid = (l + r) // 2
                self._build(2*node, l, mid)
                self._build(2*node+1, mid+1, r)
                self.cnt[node] = self.cnt[2*node] + self.cnt[2*node+1]
            
            def _push(self, node, l, r):
                if self.lazy[node] != 0:
                    mid = (l + r) // 2
                    for child, cl, cr in [(2*node, l, mid), (2*node+1, mid+1, r)]:
                        delta = self.lazy[node]
                        length = cr - cl + 1
                        # sum((d+delta)^2) = sum(d^2) + 2*delta*sum(d) + delta^2*len
                        self.sum2[child] = (self.sum2[child] + 2 * delta * self.sum1[child] + delta * delta * length) % MOD
                        self.sum1[child] = (self.sum1[child] + delta * length) % MOD
                        self.lazy[child] += delta
                    self.lazy[node] = 0
            
            def update(self, node, l, r, ql, qr, delta):
                if qr < l or ql > r:
                    return
                if ql <= l and r <= qr:
                    length = r - l + 1
                    self.sum2[node] = (self.sum2[node] + 2 * delta * self.sum1[node] + delta * delta * length) % MOD
                    self.sum1[node] = (self.sum1[node] + delta * length) % MOD
                    self.lazy[node] += delta
                    return
                self._push(node, l, r)
                mid = (l + r) // 2
                self.update(2*node, l, mid, ql, qr, delta)
                self.update(2*node+1, mid+1, r, ql, qr, delta)
                self.sum1[node] = (self.sum1[2*node] + self.sum1[2*node+1]) % MOD
                self.sum2[node] = (self.sum2[2*node] + self.sum2[2*node+1]) % MOD
            
            def query(self):
                return self.sum2[1]
        
        st = SegTree(n)
        result = 0
        
        for i in range(n):
            prev = last.get(nums[i], -1)
            # Increment d[j] by 1 for j in [prev+1, i]
            st.update(1, 0, n-1, prev + 1, i, 1)
            result = (result + st.query()) % MOD
            last[nums[i]] = i
        
        return result

Source