← 기록 / BOJ 풀이

BOJ 1699 (제곱수의 합)

오랜만에 tag: DP로 검색해서 풀어본 문제. DP 문제는 실버까지는 굉장히 재미있다. 골드부터는 관찰이 좀 어려워져서 풀이가 힘들지만.. (23년 8월 기준) PS의 빈출 영역이기 때문에 꾸준히 연습해야 한다.

이 글의 목차
  1. 풀이
  2. 1차 정답 코드
  3. 최종 정답 코드

오랜만에 tag: DP로 검색해서 풀어본 문제. DP 문제는 실버까지는 굉장히 재미있다. 골드부터는 관찰이 좀 어려워져서 풀이가 힘들지만.. (23년 8월 기준) PS의 빈출 영역이기 때문에 꾸준히 연습해야 한다.

풀이

이 문제 같은 경우 수의 범위에 주목해보면 N≤10000N \leq 10000 인데, O(N2)O(N^2)로는 TLE가 된다. 그래서 조금 머리를 굴려보면, 루프 안쪽에서 돌 때는 수의 범위가 제곱수이므로 O(NN)O(N \sqrt N)의 풀이법을 떠올릴 수 있었다.

1차 정답 코드

다만 모든 N에 대해서 최소가 되는 제곱 수의 합으로 몇번만에 업데이트 되는지에 대한 확신이 없어서 set을 섞어서 풀이했다. 나중에 생각해보니 필요가 없는 과정이었지만..

#include <bits/stdc++.h>
using namespace std;

typedef long long ll;
typedef unsigned long long ull;

int main() {
    int N; scanf("%d", &N);

    vector<int> v; v.resize(N + 1);
    fill(v.begin(), v.end(), 100000);
    vector<int> squares;
    vector<int> tmp; tmp.resize(N + 1);
    set<int> S;
    for(int i = 1; i <= N; ++i) S.insert(i);
    for(int i = 1; i * i <= N; ++i) {
        v[i*i] = 1;
        squares.push_back(i*i);
        S.erase(i*i);
    }
    while(S.size() > 0) {
        for(int i = 1; i <= N; ++i) {
            for(auto& s: squares) {
                if (i + s <= N) {
                    if (v[i + s] > v[i] + 1) {
                        v[i + s] = v[i] + 1;
                        S.erase(i + s);
                    }
                } else {
                    break;
                }
            }
        }
    }
    printf("%d\n", v[N]);

    return 0;
}

최종 정답 코드

위 코드가 깔끔하지 못해서 set을 사용한 부분을 어떻게 제거할지 고민해보니 적절한 초기값 세팅으로 set을 없앨 수 있었다. 초기값은 1=121 = 1^2임을 이용해서 세팅한다. 아래 코드가 최종 제출 코드.

#include <bits/stdc++.h>
using namespace std;

typedef long long ll;
typedef unsigned long long ull;

int main() {
    int N; scanf("%d", &N);

    int dp[100001];

    for(int i = 0; i <= N; ++i) dp[i] = i;
    for(int i = 0; i <= N; ++i) {
        for(int j = 2; j * j <= N; ++j) {
            int s = j * j;
            if (i + s <= N) {
                if (dp[i + s] > dp[i] + 1) {
                    dp[i + s] = dp[i] + 1;
                }
            } else {
                break;
            }
        }
    }
    printf("%d\n", dp[N]);

    return 0;
}