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
- Updated
Read the problem on LeetCode View on GitHub
The problem statement is LeetCode’s and stays on their site. What follows is my accepted solution.
Python
Accepted on LeetCode — runtime 6676ms, memory 58MB, accepted 2026-01-01.
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