Official

E - DNA配列の接合 / Joining of DNA Sequences Editorial by admin

Claude 4.6 Opus (Thinking)

概要

2つの文字列 \(S\), \(T\) について、\(S\) の接尾辞と \(T\) の接頭辞の最長一致(左側接合)、および \(T\) の接尾辞と \(S\) の接頭辞の最長一致(右側接合)のうち、大きい方を求める問題です。Z-function を用いて効率的に解きます。

考察

問題の言い換え

  • 左側接合: \(S\) の末尾 \(k\) 文字 = \(T\) の先頭 \(k\) 文字 となる最大の \(k\) を求める
  • 右側接合: \(T\) の末尾 \(k\) 文字 = \(S\) の先頭 \(k\) 文字 となる最大の \(k\)(ただし \(k \le L\))を求める

素朴なアプローチの問題点

\(k\) について愚直に文字列比較すると、1回の比較に \(O(k)\)、全体で \(O(L^2)\) かかり、\(L\) が最大 \(5 \times 10^5\) では TLE になります。

解決策

Z-function を利用すれば、ある文字列の各位置について「先頭との最長一致長」を \(O(N)\) で全て求められます。これを文字列連結のテクニックと組み合わせます。

アルゴリズム

Z-function とは

文字列 \(s\) に対して、\(z[i]\) = 「\(s[i:]\)\(s[0:]\) の最長共通接頭辞の長さ」を全ての \(i\) について求める配列です。線形時間で計算できます。

左側接合の求め方

\(S[L-k:L] = T[0:k]\) を満たす最大の \(k\) を求めたい。

連結文字列 \(T + \\) + S\( を作り、Z-function を計算します。位置 \)M+1+j\((\)S\( の \)j\( 文字目に対応)での Z 値は「\)T\( の先頭と \)S[j:]$ の最長共通接頭辞長」です。

\(S[j:]\) 全体(長さ \(L-j\))が \(T\) の先頭と一致する、つまり \(z[M+1+j] \ge L-j\) なら、重なり幅 \(k = L-j\) で接合可能です。\(j\)\(0\) から順に探索すれば、最初に見つかったものが最大の \(k\) です。

右側接合の求め方

\(T[M-k:M] = S[0:k]\) を満たす最大の \(k\)\(k \le L\))を求めたい。

連結文字列 \(S + \\) + T\( を作り、Z-function を計算します。位置 \)L+1+j\((\)T\( の \)j\( 文字目に対応)での Z 値は「\)S\( の先頭と \)T[j:]$ の最長共通接頭辞長」です。

\(T[j:]\) 全体(長さ \(M-j\))が \(S\) の先頭と一致する、つまり \(z[L+1+j] \ge M-j\) なら、重なり幅 \(k = M-j\) で接合可能です。\(k \le L\) の制約から \(j \ge M-L\) とし、\(j\)\(M-L\) から順に探索します。

具体例

\(S = \) 101, \(T = \) 01001 のとき:

  • 左側接合: \(S\) の末尾と \(T\) の先頭の一致を探す → 101 の末尾 01 = 01001 の先頭 01\(k=2\)
  • 右側接合: \(T\) の末尾と \(S\) の先頭の一致を探す → 01001 の末尾 1 = 101 の先頭 1\(k=1\)

答えは \(\max(2, 1) = 2\)

計算量

  • 時間計算量: \(O(L + M)\)(Z-function の計算が線形、探索も線形)
  • 空間計算量: \(O(L + M)\)(連結文字列と Z 配列の保持)

実装のポイント

  • 連結文字列にセパレータ $(元の文字列に含まれない文字)を挟むことで、Z 値が境界を超えてマッチすることを防ぐ。

  • 左側接合では \(j\) を小さい方から探すことで、最初に条件を満たした時点で最大の \(k = L - j\) が得られ、即座に break できる。

  • 右側接合では \(k \le L\) の制約に注意し、探索開始位置を \(j = \max(0, M-L)\) とする。

    ソースコード

import sys

def z_function(s):
    n = len(s)
    if n == 0:
        return []
    z = [0] * n
    z[0] = n
    l, r = 0, 0
    for i in range(1, n):
        if i < r:
            z[i] = min(r - i, z[i - l])
        while i + z[i] < n and s[z[i]] == s[i + z[i]]:
            z[i] += 1
        if i + z[i] > r:
            l, r = i, i + z[i]
    return z

def solve():
    input_data = sys.stdin.buffer.read().decode()
    tokens = input_data.split()
    L = int(tokens[0])
    M = int(tokens[1])
    S = tokens[2]
    T = tokens[3]
    
    ans = 0
    
    # Left junction: S then T
    # We need suffix of S of length k == prefix of T of length k, for max k (1 <= k <= L)
    # Build string: T + '#' + S, then use Z-function
    # We look at positions corresponding to S part, specifically the last L characters
    # For position i in the concatenated string corresponding to S[j], 
    # if z[i] >= (len(S) - j) where j is position in S... 
    # Actually, let's think differently.
    # We want: S[L-k:L] == T[0:k]
    # Concatenate: T + '$' + S
    # z[i] at position (M+1+j) for j in [0, L-1] gives the length of match between T[0:...] and S[j:...]
    # We need S[j:L] == T[0:L-j], i.e., z[M+1+j] >= L - j
    # And k = L - j, so we want to maximize k = L - j such that z[M+1+j] >= L - j
    
    concat1 = T + '$' + S
    z1 = z_function(concat1)
    
    for j in range(L):
        k = L - j  # overlap width
        idx = M + 1 + j  # position in concatenated string
        if z1[idx] >= k:
            ans = max(ans, k)
            break  # since we iterate j from 0, first match gives maximum k
    
    # Right junction: T then S
    # We need suffix of T of length k == prefix of S of length k, for max k (1 <= k <= L)
    # Concatenate: S + '$' + T
    # z[i] at position (L+1+j) for j in [0, M-1] gives match length between S[0:...] and T[j:...]
    # We need T[j:j+k] == S[0:k] where k = M - j... no wait.
    # We need T[M-k:M] == S[0:k], so j = M - k, and we need z[L+1+j] >= k = M - j
    # But k <= L, so j >= M - L
    # We want to maximize k, so minimize j. j starts from max(0, M-L).
    
    concat2 = S + '$' + T
    z2 = z_function(concat2)
    
    start_j = M - L  # since k = M - j <= L means j >= M - L
    if start_j < 0:
        start_j = 0
    
    for j in range(start_j, M):
        k = M - j  # overlap width
        if k <= 0:
            break
        idx = L + 1 + j
        if z2[idx] >= k:
            ans = max(ans, k)
            break  # first match gives maximum k since j increasing means k decreasing
    
    print(ans)

solve()

この解説は claude4.6opus-thinking によって生成されました。

posted:
last update: