Skip to content

2472. Maximum Number of Non-overlapping Palindrome Substrings

#318

Problem Link

https://leetcode.com/problems/maximum-number-of-non-overlapping-palindrome-substrings/

Problem Summary

문자열 s에서 길이가 k 이상인 팰린드롬 부분 문자열을 서로 겹치지 않게 최대 몇 개 고를 수 있는지 구하는 문제.

Solution

딱 보니 DP라서 처음에는 top-down으로 풀었다.

solve(i) = i 인덱스부터 고를 수 있는 팰린드롬의 최대 개수
solve(i) = max(solve(j + 1), s[i..j]가 길이 k 이상 팰린드롬이면 1 + solve(j + 1))   (i <= j < n)

팰린드롬 판별도 isPalindrome(l, r)로 DP를 두면 전체가 O(n^2)이 된다.

처음엔 너무 편하게 둘 다 @cache를 붙였다가 메모리 초과가 났다... isPalindrome 캐시만 200만 개라 350MB 가까이 먹는다. 일반 2차원 배열로 바꾸니 통과.

bottom-up은 팰린드롬 표를 중심에서 양쪽으로 넓혀가면서 채우고, dp[i+1]을 앞에서부터 채워주면 된다. dp 인덱스 때문에 좀 헤맸는데 거의 2배 빨라졌다.

시간복잡도는 O(n^2)

에디토리얼을 보니 그리디로도 풀린다. 자세한건 에디토리얼 참고.

앞에서부터 보면서 길이 k 이상 팰린드롬 중 가장 빨리 끝나는 걸 고르고, 그 다음부터 다시 찾으면 된다. 왜 되는지 간단하게 정리하면

  1. 최대한 많이 넣어야 되니까 빨리 끝나는 게 무조건 이득이다. 최적해의 첫 번째 팰린드롬을 가장 빨리 끝나는 걸로 바꿔도 끝나는 위치가 같거나 더 앞이라 뒤에 고른 것들과 겹치지 않는다. 즉, 개수는 그대로고 뒤에 남는 부분은 같거나 더 길어진다.
  2. 길이는 k와 k+1만 보면 된다. 팰린드롬은 양 끝을 하나씩 떼도 팰린드롬이라서 길이 k+2 이상이면 안쪽에 길이 k나 k+1짜리가 들어 있고, 그게 더 빨리 끝나니까 그걸 넣는 게 이득이다.

그래서 각 위치에서 길이 k, k+1짜리만 확인하면 O(nk)로 풀린다.

Source Code

Top-Down DP (6845ms)

class Solution:
    def maxPalindromes(self, s: str, k: int) -> int:
        n = len(s)

        paldp = [[-1 for _ in range(n)] for _ in range(n)]
        def isPalindrome(l, r):            
            if l >= r:
                return True
            if paldp[l][r] != -1:
                return paldp[l][r]

            paldp[l][r] = s[l] == s[r] and isPalindrome(l + 1, r - 1)
            return paldp[l][r]

        @cache
        def solve(i):
            if i == n:
                return 0

            ret = 0
            for j in range(i, n):
                if j - i + 1 >= k and isPalindrome(i, j):
                    ret = max(ret, 1 + solve(j + 1))
                else:
                    ret = max(ret, solve(j + 1))
            return ret

        return solve(0)

Bottom-Up DP (3766ms)

class Solution:
    def maxPalindromes(self, s: str, k: int) -> int:
        n = len(s)

        paldp = [[False for _ in range(n)] for _ in range(n)]
        for i in range(n):
            l, r = i, i
            while l >= 0 and r < n and s[l] == s[r]:
                paldp[l][r] = True
                l -= 1
                r += 1
            
            l, r = i, i + 1
            while l >= 0 and r < n and s[l] == s[r]:
                paldp[l][r] = True
                l -= 1
                r += 1

        dp = [0 for _ in range(n+1)]
        for i in range(n):
            dp[i+1] = dp[i]
            for j in range(0, i+1):
                if i - j + 1 >= k and paldp[j][i]:
                    dp[i+1] = max(dp[i+1], dp[j] + 1)
                else:
                    dp[i+1] = max(dp[i+1], dp[j])

        return dp[n]