公式

E - 図書館の貸出管理 / Library Loan Management 解説 by admin

Claude 4.6 Opus (Thinking)

概要

区間 \([L, R]\) 内の本のうち、返却予定日 \(D_i\) が指定日付 \(T\) 以下であるものの個数を高速に数える問題です。個数が \(K\) を超える場合は \(-1\) を出力します。

考察

問題の本質

各クエリは「配列 \(D\) の区間 \([L, R]\) の中で、値が \(T\) 以下である要素の個数を求めよ」という二次元的なカウントクエリに帰着されます。

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

素朴に各クエリについて \(L\) から \(R\) まで走査すると、1クエリあたり \(O(N)\)、全体で \(O(NQ)\) となり、\(N, Q \leq 10^5\) のとき最大 \(10^{10}\) 回の操作が必要になりTLEします。

解決の方針

区間に対して「ある値以下の要素数」を高速に求めるには、マージソートツリー(Merge Sort Tree)が有効です。これはセグメント木の各ノードに、対応する区間の要素をソート済み配列として保持するデータ構造です。クエリ時に二分探索を使うことで、効率的にカウントできます。

アルゴリズム

マージソートツリーの構築

  1. サイズ \(N\) の配列 \(D\) に対して、セグメント木を構築します。葉ノード \(i\) には \([D_i]\)(要素1つのリスト)を格納します。
  2. 内部ノードには、左右の子ノードのソート済み配列をマージした結果を格納します。これはマージソートと同じ操作です。

例えば \(D = [5, 3, 8, 1]\) の場合:

        [1, 3, 5, 8]       ← ルート(区間全体)
       /            \
   [3, 5]          [1, 8]   ← 内部ノード
   /    \           /    \
 [5]   [3]        [8]   [1]  ← 葉ノード

クエリの処理

区間 \([L, R]\)\(D_i \leq T\) を満たす要素数を求めるには:

  1. セグメント木上で区間 \([L, R]\)\(O(\log N)\) 個のノードに分解します。
  2. 各ノードのソート済み配列に対し、bisect_right(配列, T) を呼び出して \(T\) 以下の要素数を求めます。
  3. それらの合計が答え \(c\) です。

bisect_right は二分探索なので各ノードで \(O(\log N)\) かかり、ノード数が \(O(\log N)\) 個なので、1クエリあたり \(O(\log^2 N)\) です。

最終判定

求めた \(c\)\(K\) 以下なら \(c\) を、\(K\) を超えるなら \(-1\) を出力します。

計算量

  • 構築の時間計算量: \(O(N \log N)\)(各要素はセグメント木の深さ分 \(O(\log N)\) 回マージに参加する)
  • 1クエリの時間計算量: \(O(\log^2 N)\)\(O(\log N)\) 個のノード × 各ノードで二分探索 \(O(\log N)\)
  • 全体の時間計算量: \(O(N \log N + Q \log^2 N)\)
  • 空間計算量: \(O(N \log N)\)(セグメント木全体で保持するソート済み配列の総要素数)

\(N = Q = 10^5\) のとき、\(\log_2 N \approx 17\) なので十分高速です。

実装のポイント

  • セグメント木のサイズは \(N\) 以上の最小の2冪にすることで、インデックス計算を簡潔にしています。範囲外の葉には空リスト [] を入れます。

  • マージ操作はPythonの sorted ではなく、手動で2つのソート済みリストをマージすることで定数倍を改善しています。

  • 入力は sys.stdin.buffer.read() で一括読み込みし、高速化しています。

  • 問題の本の番号は1-indexedなので、クエリ時に L-1, R-1 として0-indexedに変換しています。

    ソースコード

import sys
from bisect import bisect_right

def main():
    input_data = sys.stdin.buffer.read().split()
    idx = 0
    N = int(input_data[idx]); idx += 1
    K = int(input_data[idx]); idx += 1
    Q = int(input_data[idx]); idx += 1
    
    D = [int(input_data[idx + i]) for i in range(N)]
    idx += N
    
    # Build a merge sort tree (persistent segment tree alternative)
    # For each node in a segment tree over [0, N-1], store sorted list of D values
    # Then for query [L, R] with threshold T, count elements <= T using bisect
    
    size = 1
    while size < N:
        size *= 2
    
    # Build merge sort tree
    tree = [None] * (2 * size)
    for i in range(N):
        tree[size + i] = [D[i]]
    for i in range(N, size):
        tree[size + i] = []
    
    for i in range(size - 1, 0, -1):
        # Merge two sorted lists
        left = tree[2 * i]
        right = tree[2 * i + 1]
        if not left:
            tree[i] = right
        elif not right:
            tree[i] = left
        else:
            # Merge
            merged = []
            a, b = 0, 0
            la, lb = len(left), len(right)
            while a < la and b < lb:
                if left[a] <= right[b]:
                    merged.append(left[a])
                    a += 1
                else:
                    merged.append(right[b])
                    b += 1
            if a < la:
                merged.extend(left[a:])
            else:
                merged.extend(right[b:])
            tree[i] = merged
    
    # Query: count elements <= T in D[L..R] (0-indexed)
    def query(l, r, T):
        # l, r are 0-indexed, inclusive
        res = 0
        l += size
        r += size + 1
        while l < r:
            if l & 1:
                res += bisect_right(tree[l], T)
                l += 1
            if r & 1:
                r -= 1
                res += bisect_right(tree[r], T)
            l >>= 1
            r >>= 1
        return res
    
    out = []
    for _ in range(Q):
        L = int(input_data[idx]); idx += 1
        R = int(input_data[idx]); idx += 1
        T = int(input_data[idx]); idx += 1
        
        c = query(L - 1, R - 1, T)
        if c > K:
            out.append("-1")
        else:
            out.append(str(c))
    
    sys.stdout.write("\n".join(out) + "\n")

main()

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

投稿日時:
最終更新: