公式

N - 株価の補正 / Stock Price Correction 解説 by admin

GPT 5.2 High

概要

「補正後の株価が毎日厳密に増加する」制約の下で、各日を整数に変更するコスト \(\sum |H_i-H'_i|\) を最小化する問題です。
これは変形すると「数列を単調非減少にする L1(絶対値)最小の等式回帰(isotonic regression)」になります。

考察

1. 「厳密に増加」を扱いやすくする変形

条件 \(H'_1 < H'_2 < \cdots < H'_N\) は、各差が少なくとも 1 あることを意味します。
ここで - \(A_i = H'_i - i\) - \(X_i = H_i - i\) と置くと、

  • \(H'_i < H'_{i+1}\)
    \(\Leftrightarrow H'_i - i \le H'_{i+1} - (i+1)\)
    \(\Leftrightarrow A_i \le A_{i+1}\)

つまり 「厳密増加」→「\(A_i\) が単調非減少」 に変換できます。

さらにコストは [ |H_i - H’_i| = |(H_i-i) - (H’_i-i)| = |X_i - Ai| ] なので、目的は [ \min \sum{i=1}^{N} |X_i - A_i| \quad \text{s.t. } A_1 \le A_2 \le \cdots \le A_N,\; A_i \in \mathbb{Z} ] となります。

  • \(X_i\) は整数なので、最適な \(A_i\) も(後述の median を使うことで)整数で達成でき、整数制約は自然に満たせます。

2. 素朴案が難しい理由

例えば DP を考えると、\(A_i\) の値域が最大 \(10^9\) 規模で現実的ではありません。
また、「単調性を壊す部分を直す」を愚直に繰り返すと、どこをいくら直すべきか(全体最適)が絡み、効率も保証できません。

3. 解決の方向性:L1 の単調回帰は「ブロックの中央値」

単調非減少にする最適解は、添字をいくつかの連続区間(ブロック)に分けて、 - 各ブロック内は同じ値 \(A_i = c\) - その \(c\) はブロック内の \(X_i\)中央値(median) にすると \(\sum |X_i - c|\) が最小になります(絶対値和の最小化の基本性質)。

したがって、 - ブロックごとに中央値を保ち - 隣り合うブロックの中央値が単調性を壊す(左 > 右)なら結合する という方針が最適です。これは PAVA(Pool Adjacent Violators Algorithm) と呼ばれる典型手法です。

アルゴリズム

全体(PAVA:中央値版)

  1. \(X_i = H_i - i\) を作る。
  2. \(X_i\) を要素 1 個のブロックとして左から順に追加する。
  3. 追加後、末尾 2 ブロックについて
    • もし「左ブロックの中央値 > 右ブロックの中央値」なら単調性違反なので、2 ブロックをマージする
    • これを違反がなくなるまで繰り返す
  4. 最終的にできた各ブロックについて、コスト(中央値からの絶対値和)を足し合わせる。

ブロック内部:中央値とコストを高速に管理

ブロック内で以下を動的に扱います: - 要素追加 - 中央値取得 - \(\sum |x - \text{median}|\)(コスト)計算

中央値を高速に保つため、2 ヒープを使います: - low:中央値以下の集合(最大ヒープ相当、コードでは符号反転して最小ヒープ) - high:中央値より大きい集合(最小ヒープ) を常に - len(low) == len(high) または len(low) == len(high)+1 となるように調整すると、中央値は常に low の最大(先頭)です。

さらに sum_low, sum_high を持つことで、中央値 \(m\) に対し [ \sum{x \in low} (m-x) + \sum{x \in high} (x-m) ] を [ m\cdot |low| - sum_low + sum_high - m\cdot |high| ] で \(O(1)\) 計算できます。

マージの高速化(小さい方を大きい方へ)

ブロック同士をマージするとき、片方の全要素をもう片方へ add していきます。
このとき常に 要素数の小さいブロックを大きいブロックへ吸収(union by size)すると、各要素が「小さい側」として移動する回数は高々 \(O(\log N)\) 回になり、全体で高速になります。

計算量

  • 時間計算量: \(O(N \log N)\)
    (各要素がヒープ操作でならされて \(O(\log N)\) 回程度しか動かないため)
  • 空間計算量: \(O(N)\)
    (全要素をヒープに保持)

実装のポイント

  • 「厳密増加」を \(X_i=H_i-i\) で「非減少」に変換するのが核心です。

  • L1(絶対値)なので、ブロック代表値は平均ではなく 中央値 になります(ここを間違えると WA)。

  • 2 ヒープ+部分和で「中央値」と「コスト」を高速に管理しています。

  • ブロックのマージは必ず 小→大 にして、最悪計算量が膨れないようにします。

    ソースコード

import sys
import heapq

class Block:
    __slots__ = ("low", "high", "sum_low", "sum_high")

    def __init__(self):
        self.low = []   # max-heap via negatives
        self.high = []  # min-heap
        self.sum_low = 0
        self.sum_high = 0

    def size(self):
        return len(self.low) + len(self.high)

    def median(self):
        return -self.low[0]

    def add(self, x: int):
        if not self.low or x <= -self.low[0]:
            heapq.heappush(self.low, -x)
            self.sum_low += x
        else:
            heapq.heappush(self.high, x)
            self.sum_high += x

        if len(self.low) < len(self.high):
            y = heapq.heappop(self.high)
            self.sum_high -= y
            heapq.heappush(self.low, -y)
            self.sum_low += y
        elif len(self.low) > len(self.high) + 1:
            y = -heapq.heappop(self.low)
            self.sum_low -= y
            heapq.heappush(self.high, y)
            self.sum_high += y

    def merge(self, other: "Block"):
        for v in other.low:
            self.add(-v)
        for v in other.high:
            self.add(v)

    def cost(self):
        m = -self.low[0]
        return m * len(self.low) - self.sum_low + self.sum_high - m * len(self.high)


def main():
    data = list(map(int, sys.stdin.buffer.read().split()))
    if not data:
        return
    N = data[0]
    H = data[1:]

    blocks = []
    for i, h in enumerate(H, start=1):
        x = h - i
        b = Block()
        b.add(x)
        blocks.append(b)

        while len(blocks) >= 2 and blocks[-2].median() > blocks[-1].median():
            b2 = blocks.pop()
            b1 = blocks.pop()
            if b1.size() < b2.size():
                b1, b2 = b2, b1
            b1.merge(b2)
            blocks.append(b1)

    ans = sum(b.cost() for b in blocks)
    print(ans)

if __name__ == "__main__":
    main()

この解説は gpt-5.2-high によって生成されました。

投稿日時:
最終更新: