---
id: "imported/blog/history/boj/boj-1699"
title: "BOJ 1699 (제곱수의 합)"
description: "오랜만에 tag: DP로 검색해서 풀어본 문제. DP 문제는 실버까지는 굉장히 재미있다. 골드부터는 관찰이 좀 어려워져서 풀이가 힘들지만.. (23년 8월 기준) PS의 빈출 영역이기 때문에 꾸준히 연습해야 한다."
kind: "record"
published: "2023-08-16T00:00:00.000Z"
tags: ["PS","BOJ","DP"]
url: "https://www.readiz.com/blog/history/boj/boj-1699/"
markdownUrl: "https://www.readiz.com/blog/history/boj/boj-1699/index.md"
---

# BOJ 1699 (제곱수의 합)

- 문제 제목: 제곱수의 합
- 문제 링크: <https://www.acmicpc.net/problem/1699>

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

## 풀이

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

### 1차 정답 코드

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

```cpp
#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 = 1^2$임을 이용해서 세팅한다. 아래 코드가 최종 제출 코드.

```cpp
#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;
}
```
