Day 034-Q2 — 三次元最長増加部分列 (3D LIS + CDQ分割統治)

2026-05-17 赤色 Master / Phase 8+ ★★★★★★★★★ CDQ + BIT

問題

$N$ 個の3次元点 $(x_i, y_i, z_i)$ が与えられる。以下を満たす最長の部分列の長さを求めよ。

部分列 $i_1 < i_2 < \ldots < i_k$ が「3次元単調増加」とは、すべての $j$ で $$x_{i_j} < x_{i_{j+1}}, \quad y_{i_j} < y_{i_{j+1}}, \quad z_{i_j} < z_{i_{j+1}}$$ を満たすこと。

制約

$1 \le N \le 10^5$
$1 \le x_i, y_i, z_i \le 10^9$
全ての点は互いに異なる

入出力例

入力例 1

5
1 3 2
2 1 3
3 2 1
4 4 4
5 5 5

出力: 3   → $(1,3,2)\to(4,4,4)\to(5,5,5)$

入力例 2

6
1 1 1
2 2 2
3 3 3
4 4 4
5 5 5
6 6 6

出力: 6   → 全点が条件を満たす

概念図: CDQ 分割統治の動作

$x$ ソート済み配列を半分に分割し、「左→右の影響」のみを $y$ マージ + $z$ BIT で処理する。

x sorted: [p1, p2, p3, p4, p5, p6, p7, p8] (N要素) LEFT [p1..p4] RIGHT [p5..p8] merge: left → right の影響 (y マージ + z BIT) 右点 p の dp[p] = max(dp[p], BIT.query(z < p.z) + 1) cdq(left) を先に cdq(right) を後で 順序: cdq(left)cross-mergecdq(right) 全体: $O(\log N)$ 段 × 各段 $O(N \log N)$ = $O(N \log^2 N)$

ヒント (段階的開示)

ヒント1: 方向性
1次元LISは $O(N \log N)$、2次元は BIT + 座標圧縮で $O(N \log^2 N)$。3次元は CDQ分割統治 (陳丹琦分治) で $O(N \log^2 N)$ 達成。
ヒント2: アプローチ
  1. $x$ でソートして1次元削減
  2. 配列を半分に分割し、cdq(left) → 「左→右の影響」 → cdq(right)
  3. マージ段階で $y$ ソート + $z$ 軸 BIT で DP 最大値を伝播
ヒント3: 擬似コード
def cdq(l, r):
    if r - l <= 1: return
    m = (l + r) // 2
    cdq(l, m)
    # 左→右の影響 (x は分割で保証 / y マージ / z BIT)
    merge_and_update_via_z_BIT(l, m, r)
    cdq(m, r)

模範解答

O(N²) 明快版
O(N log² N) CDQ完全版
import sys
from sys import stdin

def main():
    data = stdin.read().split()
    idx = 0
    N = int(data[idx]); idx += 1
    pts = []
    for i in range(N):
        x, y, z = int(data[idx]), int(data[idx+1]), int(data[idx+2])
        idx += 3
        pts.append((x, y, z))

    order = sorted(range(N), key=lambda i: pts[i])
    pts_s = [pts[order[i]] for i in range(N)]
    dp = [1] * N
    for i in range(N):
        for j in range(i):
            if pts_s[j][0] < pts_s[i][0] and pts_s[j][1] < pts_s[i][1] and pts_s[j][2] < pts_s[i][2]:
                if dp[j] + 1 > dp[i]:
                    dp[i] = dp[j] + 1
    print(max(dp))

main()
import sys
from sys import stdin

def main():
    data = stdin.read().split()
    idx = 0
    N = int(data[idx]); idx += 1
    pts = []
    for i in range(N):
        x, y, z = int(data[idx]), int(data[idx+1]), int(data[idx+2])
        idx += 3
        pts.append([x, y, z])

    # z 座標圧縮
    zs = sorted(set(p[2] for p in pts))
    zmap = {v: i+1 for i, v in enumerate(zs)}
    M = len(zs)

    # x でソート (同 x は y 昇順)
    order = sorted(range(N), key=lambda i: (pts[i][0], pts[i][1], pts[i][2]))
    pts_s = [pts[order[i]] for i in range(N)]
    dp = [1] * N

    bit = [0] * (M + 2)
    modified = []

    def bit_max(i):
        res = 0
        while i > 0:
            if bit[i] > res: res = bit[i]
            i -= i & -i
        return res

    def bit_upd(i, v):
        while i <= M:
            if bit[i] < v:
                bit[i] = v
                modified.append(i)
            i += i & -i

    def bit_clear_all():
        for i in modified:
            bit[i] = 0
        modified.clear()

    def cdq(l, r):
        if r - l <= 1: return
        m = (l + r) // 2
        cdq(l, m)

        left_y  = sorted(range(l, m), key=lambda i: pts_s[i][1])
        right_y = sorted(range(m, r), key=lambda i: pts_s[i][1])

        li = 0
        for ri in right_y:
            ry = pts_s[ri][1]
            while li < len(left_y) and pts_s[left_y[li]][1] < ry:
                lx_idx = left_y[li]
                bit_upd(zmap[pts_s[lx_idx][2]], dp[lx_idx])
                li += 1
            rz = zmap[pts_s[ri][2]]
            prev = bit_max(rz - 1)
            if prev + 1 > dp[ri]:
                dp[ri] = prev + 1

        bit_clear_all()
        cdq(m, r)

    cdq(0, N)
    print(max(dp))

main()

Step-by-Step 解説

1次元削減
3次元の条件 $x < x', y < y', z < z'$ のうち $x$ を「ソート順 = インデックス順」で吸収し、実質2次元 DP に落とす。
2CDQ 分割統治の枠組み
cdq(l, r):
  1. cdq(l, mid) を先に処理 (左の DP を確定)
  2. 「左 → 右」の影響を $y$ マージ + $z$ BIT で計算
  3. cdq(mid, r) を最後に処理
3$y$ マージ + $z$ BIT
左半分の点を $y$ 昇順で処理し、$y$ < (右点の $y$) を満たすものを BIT に投入。右点ごとに $z < r_z$ の最大 DP を取得し、$+1$ で更新。

計算量

CDQ 段数: $O(\log N)$
各段のマージ: $O(N \log N)$ (ソート + BIT)
合計: $O(N \log^2 N)$   ←  $N=10^5$ で約 $3 \times 10^7$

よくあるミス

ミス原因正しい書き方
同一 $x$ の点が影響厳密増加を無視同一 $x$ では「左→右」の影響を入れない (ソート時に同 $x$ を慎重に)
BIT クリア忘れBIT を使い回し各マージ後に modified を辿ってクリア
cdq の順序ミス右を先に処理cdq(left) → cross-merge → cdq(right) の順
$y$ 同値の扱い誤って影響厳密不等号 pts_s[left[li]][1] < ry

次のステップ

  • 発展: $k$ 次元 LIS への拡張 → $O(N \log^k N)$
  • 応用: 3D partial order counting (3次元偏順序数え上げ)

自己評価

自分の回答

気づき・メモ