Official

B - プレイリストの最適化 / Playlist Optimization Editorial by admin

gemini-3.5-flash-high

概要

曲のテンポ値を並び替えて、テンポの急激な変化(差が \(D\) を超える変化)による「違和感」の総和を最小化する問題です。

考察

1. 昇順にソートして再生する場合を考える

まず、曲をテンポ値の小さい順(昇順)に並び替えて再生することを考えてみましょう。 テンポ値をソートした配列を \(A'_1, A'_2, \ldots, A'_N\) とします。

この順で再生するとき、各曲 \(A'_j\)\(j \geq 2\))の違和感スコアは以下のようになります:

  • 直前の曲との差が \(D\) 以下(\(A'_j - A'_{j-1} \leq D\))のとき: 直前の曲が「似ている曲」にあたるため、過去に似た曲が存在することになり、違和感スコアは \(0\) になります。
  • 直前の曲との差が \(D\) より大きい(\(A'_j - A'_{j-1} > D\))のとき: 昇順に再生しているため、過去に再生された曲はすべて \(A'_j\) よりさらにテンポが小さい曲(\(A'_{j-1}\) 以下)です。したがって、過去のどの曲とも差が \(D\) より大きくなり、似ている曲は存在しません。 このとき、違和感スコアは直前の曲との差である \(A'_j - A'_{j-1}\) になります。

つまり、昇順に再生した場合、違和感の総和は「隣り合う曲のテンポの差が \(D\) を超える部分の差の合計」になります。

2. これが最小値になる理由(直感的な理解)

テンポの差が \(D\) 以下の曲同士を線で結ぶと、いくつかの「グループ」に分けることができます。 グループとグループの間には、幅が \(D\) より大きい「ギャップ(隙間)」が存在します。

どのような順番で曲を再生したとしても、最終的にはすべての曲を再生しなければなりません。 これは、数軸上で最も低いテンポのグループから、最も高いテンポのグループまで移動することを意味します。

異なるグループへ移動するとき(ギャップをまたぐとき)、過去にその移動先(またはそれ以上にテンポが高いグループ)の曲が一度も再生されていないタイミングが必ず発生します。その瞬間、ギャップを飛び越えるための違和感(=ギャップの幅以上のスコア)がどうしても発生してしまいます。

したがって、違和感の総和は「すべてのギャップの幅の合計」未満にすることはできません。 昇順にソートして再生する戦略は、この最小限の違和感(ギャップの幅の合計)をぴったり達成できるため、これが最適な再生順序となります。

アルゴリズム

  1. 与えられたテンポ値の配列 \(A\) を昇順にソートします。
  2. 答えを表す変数 ans\(0\) で初期化します。
  3. 隣り合う要素の差 \(A[i+1] - A[i]\) を順に計算します。
  4. その差が \(D\) より大きい場合のみ、その差を ans に加算します。
  5. 最終的な ans を出力します。

計算量

  • 時間計算量: \(O(N \log N)\) \(N\) 個の要素のソートに \(O(N \log N)\) かかります。その後のループは \(O(N)\) なので、全体のボトルネックはソートになり、実行時間制限に十分間に合います。
  • 空間計算量: \(O(N)\) 入力を保持するための配列 \(A\) のメモリとして \(O(N)\) の空間を使用します。

実装のポイント

  • 高速な入出力: \(N\) が最大 \(10^6\) と非常に大きいため、Pythonでは sys.stdin.read().split() を用いて入力を一括で取得することで、実行時間を大幅に短縮しています。

    ソースコード

import sys

def solve():
    # 入力を一括で読み込む
    input_data = sys.stdin.read().split()
    if not input_data:
        return
    
    N = int(input_data[0])
    D = int(input_data[1])
    A = [int(x) for x in input_data[2:]]
    
    # テンポ値を昇順にソート
    A.sort()
    
    ans = 0
    # 隣接する要素の差が D より大きい場合、その差を足し合わせる
    for i in range(N - 1):
        diff = A[i+1] - A[i]
        if diff > D:
            ans += diff
            
    print(ans)

if __name__ == '__main__':
    solve()

この解説は gemini-3.5-flash-high によって生成されました。

posted:
last update: