배열의 모든 수의 쌍의 곱의 합을 O(N)에 구하는 방법

https://atcoder.jp/contests/abc177/tasks/abc177_c

 

C - Sum of product of pairs

AtCoder is a programming contest site for anyone from beginners to experts. We hold weekly programming contests online.

atcoder.jp

 

 

 

크기 N인 배열 $A_{1}, A_{2}, ..., A_{N}$이 주어진다.

 

이때, 이 배열의 가능한 모든 인덱스 쌍 (i,j)에 대해 $A_{i} * A_{j}$의 합을 어떻게 구할까?

 

여기서 i < j이다.

 

가장 쉽게 생각할 수 있는 방법은 당연히 모든 인덱스 쌍을 순회해서 구하는 $O(N^{2})$방법이다.

 

n = int(input())

A = list(map(int,input().split()))

mod = 10**9+7

v = 0

for i in range(n-1):
    
    for j in range(i+1,n):
        
        v += A[i]*A[j]
        v %= mod

print(v)

 

 

하지만 N이 $10^{5}$ 이상 매우 큰 수라면 시간 초과에 걸린다

 

이를 해결하기 위한 방법이 다음과 같다.

 

구하고자 하는 값을 그냥 먼저 전개해보는 것이다.

 

예를 들어 N = 5라고 한다면,

 

A[0]*A[1] + A[0]*A[2] + A[0]*A[3] + A[0]*A[4]

 

+ A[1]*A[2] + A[1]*A[3] + A[1]*A[4]

 

+ A[2]*A[3] + A[2]*A[4]

 

+A[3]*A[4]

 

이걸 보기 좋게 써보면

 

A[0]*A[1]

 

A[0]*A[2] + A[1]*A[2]

 

A[0]*A[3] + A[1]*A[3] + A[2]*A[3]

                                 

A[0]*A[4] + A[1]*A[4] + A[2]*A[4] + A[3]*A[4]

 

이걸 묶어보면...

 

A[0]*A[1]

 

(A[0]+A[1])*A[2]

 

(A[0]+A[1]+ A[2])*A[3]

                                 

(A[0]+A[1]+A[2]+A[3])*A[4]

 

따라서 i = 1,2,3,4...,N-1에 대하여 A[i] * (A[0]+A[1]+A[2]+...+A[i-1])의 합을 구하면 된다.

 

(A[0]+A[1]+A[2]+...+A[i-1])은 누적합의 원리를 이용하면 쉽게 구할 수 있다.

 

따라서 이 문제는 O(N)에 해결 가능하다.

 

n = int(input())

A = list(map(int,input().split()))

mod = 10**9+7

v = 0
prefix = A[0]

for i in range(1,n):
    
    v += (prefix*A[i])
    
    v %= mod

    prefix += A[i]
    prefix %= mod

print(v)

 

 

 

또 다른 방법은 n개 항의 합의 제곱을 전개하면 나오는 공식을 이용하는 것이다.

 

$$(A_{1} + A_{2} + ... + A_{N})^{2} = A_{1}^{2} + A_{2}^{2} + ... + A_{N}^{2} + 2(A_{1}A_{2} + A_{1}A_{3} + ... + A_{N-1}A_{N})$$

 

 

놀랍게도 우변의 두번째 항 $A_{1}A_{2} + A_{1}A_{3} + ... + A_{N-1}A_{N}$이 문제에서 구하고자 하는 값이다.

 

$$A_{1}A_{2} + A_{1}A_{3} + ... + A_{N-1}A_{N} = \frac{(A_{1}+A_{2}+...+A_{N})^{2}  - (A_{1}^{2} + A_{2}^{2} + ... + A_{N}^{2})}{2}$$

 

다행히도 $(A_{1}+A_{2}+...+A_{N})^{2}$과 $A_{1}^{2} + A_{2}^{2} + ... + A_{N}^{2}$은 O(N)에 구할 수 있다.

 

따라서 이 문제는 O(N)에 해결할 수 있다.

 

n = int(input())

A = list(map(int,input().split()))

mod = 10**9+7

S = sum(A) % mod
S *= S
S %=mod

V = 0

for i in range(n):
    
    V += (A[i]**2)
    V %= mod

P = (S - V) * pow(2,mod-2,mod)

print(P % mod)

 

 

여기서 알 수 있는 사실은, 배열의 모든 수가 정수라면...

 

$(A_{1}+A_{2}+...+A_{N})^{2} - (A_{1}^{2} + ... + A_{N}^{2})$이 2의 배수라는 사실이다.

 

곱의 합이 정수니까, 2로 나누어 떨어지지 않으면 식이 모순이잖아...

 

 

그리고 보통 배열의 원소가 매우 큰 수라면 mod로 나눈 값을 구하는 문제가 나오는데

 

2로 나눈 값의 mod는 2의 모듈로 역원을 곱해야한다.

 

https://deepdata.tistory.com/577

 

모듈로 연산에서 나눗셈을 하는 방법(모듈로 곱셈의 역원 구하기)

1. 합동식에서 기본적으로 알아야하는 성질 1-1) 양변에 어떤 정수에 대한 덧셈이나 뺄셈을 하더라도 상관없다. $a \equiv b (mod p)$이고 $c \equiv d (mod p)$이면, $$a \pm c \equiv b \pm d (mod p)$$ 그러므로, c = d

deepdata.tistory.com

 

 

$10^{9}+7$이 소수이기 때문에 페르마의 소정리를 이용하면 pow(2,mod-2,mod)로 쉽게 구할 수 있다

 

728x90