Computer >> 컴퓨터 >  >> 프로그래밍 >> C++

C++에서 함수 Y = (X⁶ + X² + 9894845) % 971의 값 구하기

다음과 같은 함수가 주어졌다고 가정해 봅시다.

f(x) = (x6 + x2 + 9894845) % 971

주어진 x 값에 대해 f(x)의 값을 구하는 것이 목표입니다. 예를 들어 입력이 5라면 결과는 469가 됩니다.

접근 방법

x가 커지면 x6처럼 거듭제곱한 값이 급격히 커져 자료형의 표현 범위를 벗어나는 오버플로우가 발생할 수 있고, 연산 속도도 느려집니다. 이를 해결하려면 모듈러 거듭제곱(modular exponentiation), 즉 빠른 거듭제곱(fast exponentiation) 기법을 사용해야 합니다. 이 방법은 분할 정복 원리를 이용해 (base^exponent) % modulus를 O(log n) 시간 복잡도로 계산합니다.

여기에 한 가지 최적화를 더할 수 있습니다. 9894845 % 971 = 355이므로 상수항을 미리 355로 줄여 두면 계산이 더 간단해집니다.

풀이 단계

  • 밑(base), 지수(exponent), 모듈러스(modulus)를 인자로 받는 함수 power_mod()를 정의합니다.
  • base := base mod modulus로 초기화합니다.
  • result := 1로 초기화합니다.
  • exponent > 0인 동안 다음 과정을 반복합니다.
    • exponent가 홀수이면 result := (result × base) mod modulus
    • base := (base × base) mod modulus
    • exponent := exponent ÷ 2
  • 반복이 끝나면 result를 반환합니다.
  • 메인 함수에서는 다음 식을 계산해 최종 답을 구합니다.
    ((power_mod(n, 6, m) + power_mod(n, 2, m)) % m + 355) % m

예제 코드

아래 구현을 통해 더 자세히 이해해 보겠습니다.

#include <bits/stdc++.h>
using namespace std;
typedef long long int lli;

lli power_mod(lli base, lli exponent, lli modulus) {
    base %= modulus;
    lli result = 1;
    while (exponent > 0) {
        if (exponent & 1)
            result = (result * base) % modulus;
        base = (base * base) % modulus;
        exponent >>= 1;
    }
    return result;
}

int main() {
    lli n, m = 971;
    cin >> n;
    cout << (((power_mod(n, 6, m) + power_mod(n, 2, m)) % m + 355) % m);
    return 0;
}

입력

84562

출력

140

동작 원리 정리

power_mod()는 지수를 절반씩 줄여 가면서 밑을 계속 제곱하는 방식으로 동작합니다. 지수가 홀수일 때만 현재 밑을 결과에 곱해 주고, 매 단계마다 모듈러스로 나머지를 취하기 때문에 중간값이 커지지 않습니다. 덕분에 지수가 작든 크든 항상 빠른 시간 안에 정확한 결과를 얻을 수 있습니다.