문제 개요
m × n 크기의 행렬 mat와 하나의 정수 threshold(임계값)가 주어졌을 때, 모든 원소의 합이 임계값 이하인 정사각형 중에서 가장 큰 변의 길이를 구하는 문제입니다. 만약 조건을 만족하는 정사각형이 하나도 존재하지 않는다면 0을 반환해야 합니다.
예를 들어 입력 행렬이 다음과 같다고 가정해 보겠습니다.
| 1 | 1 | 3 | 2 | 4 | 3 | 2 |
| 1 | 1 | 3 | 2 | 4 | 3 | 2 |
| 1 | 1 | 3 | 2 | 4 | 3 | 2 |
여기서 임계값이 4라면, 아래 표에서 초록색으로 표시한 부분처럼 변의 길이가 2인 정사각형(합 = 1+1+1+1 = 4)을 찾을 수 있습니다. 이러한 정사각형이 총 두 개 존재하므로 그중 최댓값인 2가 정답이 됩니다.
| 1 | 1 | 3 | 2 | 4 | 3 | 2 |
| 1 | 1 | 3 | 2 | 4 | 3 | 2 |
| 1 | 1 | 3 | 2 | 4 | 3 | 2 |
해결 접근 방법
이 문제는 2차원 누적 합(Prefix Sum)과 이진 탐색(Binary Search)을 결합하면 효율적으로 해결할 수 있습니다. 행렬을 미리 누적 합 형태로 변환해 두면 임의의 정사각형 영역의 합을 O(1) 시간 안에 계산할 수 있고, 변의 길이 후보에 대해 이진 탐색을 수행해 최댓값을 빠르게 찾을 수 있습니다.
구체적인 해결 단계는 다음과 같습니다.
- 변의 길이가 x인 정사각형 중 합이 th 이하인 것이 존재하는지 확인하는 함수 ok(x, mat, th)를 정의합니다.
- curr := 0으로 초기화합니다.
- r을 x − 1부터 마지막 행 인덱스까지 순회합니다.
- c를 x − 1부터 마지막 열 인덱스까지 순회합니다.
- curr := mat[r][c]
- c − x ≥ 0이면 curr에서 mat[r][c − x]를 뺍니다.
- r − x ≥ 0이면 curr에서 mat[r − x][c]를 뺍니다.
- c − x ≥ 0이고 r − x ≥ 0이면 curr에 mat[r − x][c − x]를 다시 더합니다. (포함·배제 원리로 중복 제거)
- curr ≤ th이면 true를 반환합니다.
- c를 x − 1부터 마지막 열 인덱스까지 순회합니다.
- 전체를 순회한 후에도 조건을 만족하는 정사각형이 없으면 false를 반환합니다.
- 메인 함수에서는 행렬과 임계값을 입력받아 처리합니다.
- r := 행의 개수, c := 열의 개수, low := 1, high := min(r, c), ans := 0으로 초기화합니다.
- 먼저 열 방향 누적 합을 계산합니다. (i = 1 ~ c − 1, 각 행 j에 대해 mat[j][i] += mat[j][i − 1])
- 이어서 행 방향 누적 합을 계산합니다. (i = 1 ~ r − 1, 각 열 j에 대해 mat[i][j] += mat[i − 1][j])
- low ≤ high인 동안 이진 탐색을 진행합니다.
- mid := low + (high − low) / 2
- ok(mid, mat, th)가 참이면 ans := mid, low := mid + 1 → 더 큰 변의 길이를 탐색
- 거짓이면 high := mid − 1 → 더 작은 변의 길이를 탐색
- 탐색이 종료되면 ans를 반환합니다.
복잡도 분석
- 시간 복잡도: O(min(r, c) × r × c) — 이진 탐색의 각 단계에서 ok() 함수가 전체 행렬을 한 번씩 순회하기 때문입니다.
- 공간 복잡도: O(1) — 입력 행렬 자체를 누적 합 배열로 재활용하므로 별도의 추가 공간이 필요하지 않습니다.
C++ 구현 예제
아래 코드를 통해 더 자세히 이해해 보겠습니다.
#include <bits/stdc++.h>
using namespace std;
typedef long long int lli;
class Solution {
public:
bool ok(int x, vector<vector<int>>& mat, int th){
lli current = 0;
for(int r = x - 1; r < mat.size(); r++){
for(int c = x - 1; c < mat[0].size(); c++){
current = mat[r][c];
if(c - x >= 0)current -= mat[r][c-x];
if(r - x >= 0)current -= mat[r - x][c];
if(c - x >= 0 && r - x >= 0)current += mat[r-x][c-x];
if(current <= th)return true;
}
}
return false;
}
int maxSideLength(vector<vector<int>>& mat, int th) {
int r = mat.size();
int c = mat[0].size();
int low = 1;
int high = min(r, c);
int ans = 0;
for(int i = 1; i < c; i++){
for(int j = 0; j < r; j++){
mat[j][i] += mat[j][i - 1];
}
}
for(int i = 1; i < r; i++){
for(int j = 0; j < c; j++){
mat[i][j] += mat[i - 1][j];
}
}
while(low <= high){
int mid = low + (high - low) / 2;
if(ok(mid, mat, th)){
ans = mid;
low = mid + 1;
}
else{
high = mid - 1;
}
}
return ans;
}
};
main(){
vector<vector<int>> v = {{1,1,3,2,4,3,2},{1,1,3,2,4,3,2},{1,1,3,2,4,3,2}};
Solution ob;
cout << (ob.maxSideLength(v, 4));
}입력
[[1,1,3,2,4,3,2],[1,1,3,2,4,3,2],[1,1,3,2,4,3,2]] 4
출력
2
출력 결과가 2인 이유는, 변의 길이가 2인 정사각형의 합은 정확히 4로 임계값 이하이지만, 변의 길이가 3인 정사각형의 합은 15로 임계값을 초과하기 때문입니다.