s와
s[1:]
s[2:]
...
s[len(s) - 1:]
def brute_force(s: str) -> list[int]:
N = len(s)
z_array = [0] * N
for i in range(N):
count = 0
while i + count < N and s[count] == s[i + count]:
count += 1
z_array[i] = count
return z_array
s = "ababa" 라고 해보자.s = ababa
s[0:] = ababa -> ababa와 비교 -> 5
s[1:] = baba -> ababa와 비교 -> 0
s[2:] = aba -> ababa와 비교 -> 3
s[3:] = ba -> ababa와 비교 -> 0
s[4:] = a -> ababa와 비교 -> 1
[5, 0, 3, 0, 1]
s = "aaaaaa"
s[0:] = aaaaaa -> 6번 비교
s[1:] = aaaaa -> 5번 비교
s[2:] = aaaa -> 4번 비교
s[3:] = aaa -> 3번 비교
s[4:] = aa -> 2번 비교
s[5:] = a -> 1번 비교
N + (N - 1) + (N - 2) + ... + 1
O(N^2) 이 된다.N = 10^5 같은 입력에서는 사실상 불가능하다.prefix: a b a b a ...
suffix: a b a ...
aba가 같다는 사실을 알았다면,이미 prefix와 일치한다고 확인한 구간을 기억해두고,
그 안에 들어오는 위치는 가능한 만큼 기존 값을 재사용한다.
[left, right] 라고 둔다.s = a b a b a ...
^ ^
left right
s[left:right + 1]구간은
s[0:right - left + 1]와 동일하다.
left에서 시작하는 어떤 부분 문자열이 prefix와 일치하고 있다는 뜻이다.s[0 ...] = a b a
s[left ...] = a b a
left와 right 사이에 있는 인덱스들은 처음부터 다시 비교할 필요가 없다.i가 right보다 크다면, 현재 알고 있는 일치 구간 밖에 있는 것이다. [left .... right] i
left = right = i
while right < N and s[right - left] == s[right]:
right += 1
right -= 1
z_array[i] = right - left + 1
right - left는 prefix 쪽 인덱스다.left = i 라면,s[0] vs s[i]
s[1] vs s[i + 1]
s[2] vs s[i + 2]
...
i <= right인 경우다. left i right
|----------|---------|
i는 이미 prefix와 일치한다고 검증된 구간 안에 있다.i에 대응되는 prefix 쪽 위치를 찾을 수 있다.k = i - left
[left, right] 구간은 prefix와 동일하기 때문이다.s[0] == s[left]
s[1] == s[left + 1]
s[2] == s[left + 2]
...
s[k] == s[i]
z_array[k] 값을 참고할 수 있다.remaining = right - i + 1
i부터 Z-box 끝까지 남아있는 길이다. left i right
|----------|---------|
<--------->
remaining
z_array[k]와 remaining을 비교한다.if z_array[k] < remaining:
z_array[i] = z_array[k]
z_array[k]가 Z-box 내부에서 완전히 끝난다는 뜻이다.prefix 쪽에서 이미 불일치가 확인된 위치가
현재 Z-box 내부에도 그대로 대응된다.
z_array[k]를 그대로 가져오면 된다.z_array[i] = z_array[k]
else:
left = i
while right + 1 < N and s[right - left + 1] == s[right + 1]:
right += 1
z_array[i] = right - left + 1
z_array[k]가 현재 Z-box 끝까지 닿거나, 그 밖으로 튀어나갈 가능성이 있다는 뜻이다. left i right
|----------|---------|
<--------->
remaining
z_array[k]가 remaining 이상
right 이후는 아직 비교해본 적이 없다.right + 1부터 추가 비교를 해야 한다.이미 아는 부분: i ~ right
새로 확인할 부분: right + 1부터
left = i로 갱신하고,def z_algorithm(s: str) -> list[int]:
N = len(s)
z_array = [0] * N
z_array[0] = N
left = right = 0
for i in range(1, N):
if i > right:
"""
Outside of the Z-box
"""
left = right = i
while right < N and s[right - left] == s[right]:
right += 1
right -= 1
z_array[i] = right - left + 1
else:
"""
Inside of the Z-box
"""
k = i - left
remaining = right - i + 1
if z_array[k] < remaining:
z_array[i] = z_array[k]
else:
left = i
while right + 1 < N and s[right - left + 1] == s[right + 1]:
right += 1
z_array[i] = right - left + 1
return z_array
O(N^2)처럼 보일 수 있다.right가 한 번 증가하면 다시 왼쪽으로 돌아가지 않는다는 것이다.right는 전체 알고리즘 동안 최대 N번만 증가한다.
O(N)
O(N)
O(N^2)이 된다.[left, right]를 관리한다.O(N)으로 줄일 수 있다.text 안에서 pattern이 등장하는 위치를 찾을 수 있다.예를 들어 다음과 같은 상황을 생각해보자.
pattern = "aba"
text = "abacaba"
text 안에서 pattern이 시작되는 위치는 어디인가?
aba는 다음 위치에 등장한다.text = a b a c a b a
0 1 2 3 4 5 6
pattern = aba
등장 위치: 0, 4
z-algorithm은 항상 어떤 문자열의 prefix와 suffix를 비교한다.
그러면 패턴 매칭을 하려면 어떻게 해야 할까?
간단하다.
pattern을 앞에 두고, 그 뒤에 text를 붙이면 된다.
combined = pattern + "#" + text
#는 구분자다.pattern = aba
text = abacaba
combined = aba#abacaba
# 같은 구분자를 넣을까?aba#abacaba
combined의 prefix와
combined[i:]가 얼마나 길게 일치하는가?
pattern이다.combined prefix = aba
len(pattern) 이상이라면,combined = aba#abacaba
index = 0 1 2 3 4 5 6 7 8 9 10
combined = a b a # a b a c a b a
text가 시작되는 위치는 len(pattern) + 1이다.text_start = len(pattern) + 1
text_start = 3 + 1
text_start = 4
combined = a b a # a b a c a b a
index = 0 1 2 3 4 5 6 7 8 9 10
^
text 시작
combined = aba#abacaba
z_array = [11, 0, 1, 0, 3, 0, 1, 0, 3, 0, 1]
3이다.len(pattern) = 3
index 4 -> z_array[4] = 3
index 8 -> z_array[8] = 3
"aba"가 완전히 일치한다는 뜻이다.combined[4:] = abacaba
aba 일치
combined[8:] = aba
aba 일치
그런데 우리가 원하는 것은 combined 기준 인덱스가 아니다.
원래 text 기준 인덱스가 필요하다.
combined에서 text는 4번 인덱스부터 시작했다.
따라서 다음처럼 빼주면 된다.
text_index = combined_index - text_start
combined index 4 -> text index 0
combined index 8 -> text index 4
[0, 4]
z_algorithm 함수를 그대로 사용한다.def z_algorithm(s: str) -> list[int]:
N = len(s)
z_array = [0] * N
z_array[0] = N
left = right = 0
for i in range(1, N):
if i > right:
left = right = i
while right < N and s[right - left] == s[right]:
right += 1
right -= 1
z_array[i] = right - left + 1
else:
k = i - left
remaining = right - i + 1
if z_array[k] < remaining:
z_array[i] = z_array[k]
else:
left = i
while right + 1 < N and s[right - left + 1] == s[right + 1]:
right += 1
z_array[i] = right - left + 1
return z_array
def find_pattern(text: str, pattern: str) -> list[int]:
combined = pattern + "#" + text
z_array = z_algorithm(combined)
pattern_length = len(pattern)
text_start = pattern_length + 1
result = []
for i in range(text_start, len(combined)):
if z_array[i] >= pattern_length:
result.append(i - text_start)
return result
text = "abacaba"
pattern = "aba"
print(find_pattern(text, pattern))
[0, 4]
combined = pattern + "#" + text
z_array[i] >= len(pattern)
combined[i:i + len(pattern)] == pattern
def brute_force_pattern_matching(text: str, pattern: str) -> list[int]:
result = []
for i in range(len(text) - len(pattern) + 1):
matched = True
for j in range(len(pattern)):
if text[i + j] != pattern[j]:
matched = False
break
if matched:
result.append(i)
return result
O(NM)
여기서 N은 text의 길이,
M은 pattern의 길이다.
예를 들어 다음처럼 같은 문자가 반복되는 경우를 생각해보자.
text = "aaaaaaaaaa"
pattern = "aaaaa"
각 위치마다 거의 pattern 전체를 비교해야 한다.
그래서 입력이 커지면 상당히 느려질 수 있다.
반면 z-algorithm을 사용하면 combined 문자열 하나에 대해 z-array를 한 번만 구하면 된다.
combined length = len(pattern) + 1 + len(text)
O(N + M)
def z_algorithm(s: str) -> list[int]:
N = len(s)
z_array = [0] * N
z_array[0] = N
left = right = 0
for i in range(1, N):
if i > right:
left = right = i
while right < N and s[right - left] == s[right]:
right += 1
right -= 1
z_array[i] = right - left + 1
else:
k = i - left
remaining = right - i + 1
if z_array[k] < remaining:
z_array[i] = z_array[k]
else:
left = i
while right + 1 < N and s[right - left + 1] == s[right + 1]:
right += 1
z_array[i] = right - left + 1
return z_array
def find_pattern(text: str, pattern: str) -> list[int]:
combined = pattern + "#" + text
z_array = z_algorithm(combined)
pattern_length = len(pattern)
text_start = pattern_length + 1
result = []
for i in range(text_start, len(combined)):
if z_array[i] >= pattern_length:
result.append(i - text_start)
return result