Skip to content

301. Remove Invalid Parentheses

#200

Problem link

https://leetcode.com/problems/remove-invalid-parentheses/

Problem Summary

괄호가 맞도록 최소한의 문자를 제거해서 올바른 괄호 문자열을 출력하는 문제.

Solution

처음엔 간단한 완전탐색 dfs를 했는데 시간이 너무 오래 걸림. (심지어 초반 코드는 시간 초과라 최적화를 조금 했지만 9s 넘게 걸렸다)
=> 최소한의 문자를 제거하는 것이므로 BFS 방식으로 변경

현재 상태를 level이란 set으로 두자. level에서 가능한 문자열이 있으면 리턴하고 아니라면 그 level의 문자열 중에서 모든 문자를 하나씩 제거해서 다음 level에 집어넣는 방식으로 작성하였다. (discuss 참고함)

Source Code

DFS (9405 ms)

from typing import List


class Solution:
    def removeInvalidParentheses(self, s: str) -> List[str]:
        n = len(s)

        def is_valid(st):
            stack = []

            for c in st:
                if c == '(':
                    stack.append(c)
                elif c == ')':
                    if not stack:
                        return False
                    stack.pop()

            return not stack

        res = dict()
        visited = {}

        def dfs(st, start, cnt):
            if res and cnt > min(res.keys()):
                return

            if start == len(st):
                if is_valid(st):
                    if cnt not in res:
                        res[cnt] = []
                    if st in visited:
                        return
                    res[cnt].append(st)
                    visited[st] = True
                return

            if st[start] in '()':
                dfs(st, start + 1, cnt)
                dfs(st[:start] + st[start + 1:], start, cnt + 1)
            else:
                dfs(st, start + 1, cnt)

        dfs(s, 0, 0)

        return res[min(res.keys())

BFS (192 ms)

from typing import List


class Solution:
    def removeInvalidParentheses(self, s: str) -> List[str]:
        def is_valid(st):
            stack = []

            for c in st:
                if c == '(':
                    stack.append(c)
                elif c == ')':
                    if not stack:
                        return False
                    stack.pop()

            return not stack

        level = {s}

        while True:
            res = []
            for st in level:
                if is_valid(st):
                    res.append(st)

            if res:
                return res

            new_level = set()
            for st in level:
                for i in range(len(st)):
                    if st[i] in '()':
                        new_st = st[:i] + st[i + 1:]
                        new_level.add(new_st)

            level = new_level