Computer >> 컴퓨터 >  >> 프로그래밍 >> Python

Python으로 숫자 리스트의 모든 쿼리에 대한 kpr 합계 계산하기

문제 이해하기

숫자 리스트 nums가 주어져 있다고 가정해 보겠습니다. 그리고 각 쿼리 queries[i]가 세 개의 값 [k, p, r]을 담고 있는 쿼리 리스트도 함께 제공됩니다. 각 쿼리에 대해 kpr_sum을 계산해야 하며, 그 공식은 다음과 같습니다.

$$\mathrm{kpr\_sum} = \sum_{i=P}^{R-1}\sum_{j=i+1}^{R}(K \oplus (A[i] \oplus A[j]))$$

계산 결과가 너무 커질 경우에는 10^9+7로 나눈 나머지를 반환하면 됩니다.

예를 들어 입력이 nums = [1,2,3], queries = [[1,1,3],[2,1,3]]이라면 출력은 [5, 4]가 됩니다. 첫 번째 쿼리의 경우 (1 XOR (1 XOR 2)) + (1 XOR (1 XOR 3)) + (1 XOR (2 XOR 3)) = 5이며, 두 번째 쿼리도 같은 방식으로 계산하면 4가 됩니다.

해결 접근 방법

이 문제는 각 비트 자릿수별로 누적합(프리픽스 합)을 미리 구해 두면 효율적으로 해결할 수 있습니다. 특정 비트 위치 i에 대해 구간 [P, R] 안에 해당 비트가 1인 숫자의 개수(n1)와 0인 숫자의 개수(n0)를 빠르게 파악한 뒤, K의 i번째 비트 값에 따라 조건을 만족하는 페어의 개수를 계산하는 원리입니다.

  • K의 i번째 비트가 1인 경우: (A[i] ⊕ A[j])의 i번째 비트가 0일 때 결과 비트가 1이 되므로, 두 숫자의 해당 비트가 서로 같은 경우의 수, 즉 n1개 중에서 고르는 경우와 n0개 중에서 고르는 경우를 더한 (n1×(n1−1) + n0×(n0−1))/2가 됩니다.
  • K의 i번째 비트가 0인 경우: 두 숫자의 해당 비트가 서로 달라야 하므로 n1 × n0개의 페어가 됩니다.

이를 단계별로 정리하면 다음과 같습니다.

  • m := 10^9 + 7
  • N := nums의 크기
  • q_cnt := queries의 크기
  • C := 새로운 리스트
  • res := 새로운 리스트
  • i를 0부터 19까지 반복합니다:
    • R := 0 하나만 담고 있는 배열
    • t := 0
    • nums의 각 x에 대해:
      • t := t + ((x를 i번 오른쪽 시프트한 값) AND 1)
      • R의 끝에 t를 삽입
    • C의 끝에 R을 삽입
  • j를 0부터 q_cnt까지 반복합니다:
    • (K, P, R) := queries[j]
    • d := R − P + 1
    • t := 0
    • i를 0부터 19까지 반복합니다:
      • n1 := C[i][R] − C[i][P−1]
      • n0 := d − n1
      • 만약 (K를 i번 오른쪽 시프트한 값) AND 1이 0이 아니라면:
        • x := (n1 × (n1 − 1) + n0 × (n0 − 1)) / 2의 몫
      • 그렇지 않으면:
        • x := n1 × n0
      • t := (t + (x를 i번 왼쪽 시프트한 값)) mod m
    • res의 끝에 t를 삽입
  • res 반환

예제 코드

아래 구현을 통해 더 잘 이해해 보겠습니다.

def solve(nums, queries):
    m = 10**9 + 7
    N = len(nums)
    q_cnt = len(queries)
    C = []
    res = []
    for i in range(20):
        R = [0]
        t = 0
        for x in nums:
            t += (x >> i) & 1
            R.append(t)
        C.append(R)
    for j in range(q_cnt):
        K, P, R = queries[j]
        d = R - P + 1
        t = 0
        for i in range(20):
            n1 = C[i][R] - C[i][P-1]
            n0 = d - n1
            if (K >> i) & 1:
                x = (n1 * (n1 - 1) + n0 * (n0 - 1)) >> 1
            else:
                x = n1 * n0
            t = (t + (x << i)) % m
        res.append(t)

    return res

nums = [1,2,3]
queries = [[1,1,3],[2,1,3]]
print(solve(nums, queries))

입력

[1,2,3], [[1,1,3],[2,1,3]]

출력

[5, 4]

시간 복잡도

전처리 단계는 O(N × 20), 각 쿼리 처리는 O(20)이므로 전체 시간 복잡도는 O((N + Q) × 20)입니다. 구간 내 모든 페어를 직접 계산하는 O(Q × N²) 방식보다 훨씬 효율적이므로, 리스트의 크기와 쿼리 수가 모두 클 때에도 빠르게 동작합니다.