Imports
open Finsetopen scoped BigOperators14.2. Matrix-Chain Multiplication
The legacy CLRS.Chapter15.matrixChainOpt and
CLRS.Chapter15.matrixChainSplit recursively evaluate an optimization
specification; they do not store a dynamic-programming table. This section
retains their arithmetic table-size and candidate-budget formulas.
The Execution companion supplies MatrixChainExecution.execute:
actual arrays of interval-length rows storing costs and selected splits. Each
candidate reads strictly shorter stored intervals. Its correctness theorem
identifies every stored cost and proves the selected split attains it. Table
reconstruction reads those stored splits without rerunning the recurrence.
The cell and candidate counters are accumulated by that same fill execution.
The parameter N in the execution is the largest matrix index, so indices
0..N describe N + 1 matrices. Array access, dimension lookup, and
arithmetic are primitive events; allocation/copying and arithmetic bit costs
are excluded. The formulas below alone are not execution-cost proofs; use the
companion's execute_cells and execute_candidateVisits interfaces.
Notation conventions used in this section:
-
dims: the dimension table,dims i= the number of rows of matrixAᵢ -
n: the number of matrices
namespace CLRSnamespace Chapter15Space and time of the table algorithm
The number of distinct (i, j) subproblems with 0 ≤ i ≤ j ≤ n,
an arithmetic interval-count formula.
def matrixChainSpace (n : Nat) : Nat :=
(n + 1) * (n + 2) / 2
An abstract split-count formula: for each
interval [i, j] there are j - i candidate split points.
def matrixChainTime (n : Nat) : Nat :=
(Finset.range (n + 1)).sum (fun j => (Finset.range j).sum (fun i => j - i))
The table has (n + 1)(n + 2) / 2 entries.
theorem matrixChainSpace_eq (n : Nat) :
matrixChainSpace n = (n + 1) * (n + 2) / 2 := rflQuadratic bound on the interval-count formula.
theorem matrixChainSpace_le_square (n : Nat) : matrixChainSpace n ≤ (n + 2) ^ 2 := by
unfold matrixChainSpace
calc
(n + 1) * (n + 2) / 2 ≤ (n + 1) * (n + 2) := Nat.div_le_self _ _
_ ≤ (n + 2) * (n + 2) := Nat.mul_le_mul_right _ (by omega : n + 1 ≤ n + 2)
_ = (n + 2) ^ 2 := by rw [pow_two]Cubic bound on the abstract split-count formula.
theorem matrixChainTime_le_cubic (n : Nat) : matrixChainTime n ≤ (n + 1) ^ 3 := by
unfold matrixChainTime
calc
(Finset.range (n + 1)).sum (fun j => (Finset.range j).sum (fun i => j - i))
≤ (Finset.range (n + 1)).sum (fun j => j ^ 2) := by
apply Finset.sum_le_sum
intro j hj
calc
(Finset.range j).sum (fun i => j - i) ≤ (Finset.range j).sum (fun _ => j) := by
apply Finset.sum_le_sum
intro i hi
omega
_ = j * j := by simp [Finset.sum_const, Finset.card_range]
_ = j ^ 2 := by rw [pow_two]
_ ≤ (n + 1) * n ^ 2 := by
have hbound : ∀ j ∈ Finset.range (n + 1), j ^ 2 ≤ n ^ 2 := by
intro j hj
have hjle : j ≤ n := by simpa [mem_range] using hj
exact Nat.pow_le_pow_left hjle 2
simpa [Finset.card_range, nsmul_eq_mul] using
(Finset.sum_le_card_nsmul (Finset.range (n + 1)) (fun j => j ^ 2) (n ^ 2)
(by intro j hj; exact hbound j hj))
_ ≤ (n + 1) ^ 3 := by
have h : n ^ 2 ≤ (n + 1) ^ 2 := Nat.pow_le_pow_left (Nat.le_succ n) 2
have h' : (n + 1) * n ^ 2 ≤ (n + 1) * (n + 1) ^ 2 := Nat.mul_le_mul_left (n + 1) h
rw [show (n + 1) ^ 3 = (n + 1) * (n + 1) ^ 2 by rw [pow_succ, mul_comm]]
exact h'end Chapter15end CLRSDefinitions and proofs
CLRSLean.FourthEdition.Chapter_14.Section_14_2_Matrix_Chain_Multiplication.Execution
Stored matrix-chain costs and splits
The endpoint N denotes matrices indexed 0 through N, hence N+1 matrices.
Layer l stores intervals [i,i+l]; only shorter layers are read. Every split
candidate is evaluated once, and its minimizing index is stored with its cost.
The old recursive optimum is used only in proofs.
namespace CLRS.Chapter15.MatrixChainExecutionopen DPExecutionstructure Cell where
cost : Nat
split : Nat
deriving Inhabited, Repr, DecidableEqdef candidate (dims : Nat → Nat) (rows : Array (Array Cell))
(i l k : Nat) : Nat :=
(get rows (k-i) i).cost + (get rows (i+l-(k+1)) (k+1)).cost +
dims i * dims (k+1) * dims (i+l+1)def step (dims : Nat → Nat) : Nat → Array (Array Cell) → Nat → Cell × Nat
| 0, _, i => (⟨0, i⟩, 0)
| l+1, rows, i =>
let best := minimum (candidate dims rows i (l+1)) i l
(⟨best.value, best.index⟩, best.visits)Compute and retain all intervals 0≤i≤j≤N.
def execute (dims : Nat → Nat) (N : Nat) : Table Cell :=
buildLayers (fun l => N+1-l) (step dims) (N+1)The stored split is in range and attains the independently defined optimum.
def CorrectCell (dims : Nat → Nat) (l i : Nat) (cell : Cell) : Prop :=
cell.cost = matrixChainOpt dims i (i+l) ∧
(0 < l → i ≤ cell.split ∧ cell.split < i+l ∧
cell.cost = matrixSplitCost dims (matrixChainOpt dims) i (i+l) cell.split)
private theorem candidate_correct (dims : Nat → Nat) (N l i : Nat)
(rows : Array (Array Cell))
(hprev : ∀ k, k < l → ∀ i, i < N+1-k → CorrectCell dims k i (get rows k i))
(hi : i < N+1-l) (k : Nat) (hk : i ≤ k) (hk' : k < i+l) :
candidate dims rows i l k = matrixSplitCost dims (matrixChainOpt dims) i (i+l) k := by
have hleft := (hprev (k-i) (by omega) i (by omega)).1
have hright := (hprev (i+l-(k+1)) (by omega) (k+1) (by omega)).1
have he : i + (k-i) = k := by omega
have he' : k+1+(i+l-(k+1)) = i+l := by omega
rw [he] at hleft
rw [he'] at hright
simp only [candidate, matrixSplitCost, hleft, hright]
private theorem step_correct (dims : Nat → Nat) (N l : Nat)
(rows : Array (Array Cell))
(hprev : ∀ k, k < l → ∀ i, i < N+1-k → CorrectCell dims k i (get rows k i))
(i : Nat) (hi : i < N+1-l) : CorrectCell dims l i (step dims l rows i).1 := by
cases l with
| zero =>
simp [CorrectCell, step, matrixChainOpt]
| succ l =>
have hc := candidate_correct dims N (l+1) i rows hprev hi
let f := candidate dims rows i (l+1)
have hbest : (minimum f i l).value = matrixChainOpt dims i (i+(l+1)) := by
apply minimum_value_eq
· intro k hk hk'
dsimp only [f]
rw [hc k hk (by omega)]
exact (matrixChainOpt_lowerBound dims).2 (by simp [Finset.mem_Icc]; omega)
· obtain ⟨hk, heq⟩ := matrixChainSplit_optimal dims i (i+(l+1)) (by omega)
have hb := Finset.mem_Icc.mp hk
refine ⟨matrixChainSplit dims i (i+(l+1)), hb.1, by omega, ?_⟩
exact (hc _ hb.1 (by omega)).trans heq.symm
obtain ⟨hlo, hhi, heq, _⟩ := minimum_spec f i l
change (minimum f i l).value = _ ∧
(0 < l+1 → i ≤ (minimum f i l).index ∧ (minimum f i l).index < i+(l+1) ∧
(minimum f i l).value = _)
refine ⟨hbest, fun _ => ⟨hlo, by omega, ?_⟩⟩
exact heq.trans (hc _ hlo (by omega))Every stored entry is the optimum, with a valid tight stored split.
theorem execute_get_correct (dims : Nat → Nat) (N l i : Nat)
(hl : l ≤ N) (hi : i+l ≤ N) :
CorrectCell dims l i (get (execute dims N).rows l i) := by
apply buildLayers_property (fun l => N+1-l) (step dims) (CorrectCell dims)
(fun l rows _ hp i hi => step_correct dims N l rows hp i hi) (N+1) l (by omega) i
(by omega)
A direct cost-table refinement for interval [i,j].
theorem execute_cost (dims : Nat → Nat) (N i j : Nat) (hij : i ≤ j) (hj : j ≤ N) :
(get (execute dims N).rows (j-i) i).cost = matrixChainOpt dims i j := by
have h := (execute_get_correct dims N (j-i) i (by omega) (by omega)).1
simpa [Nat.add_sub_of_le hij] using h@[simp] theorem step_visits (dims : Nat → Nat) (l : Nat) (rows : Array (Array Cell)) (i : Nat) :
(step dims l rows i).2 = l := by cases l <;> simp [step]Each interval cell is written once by the actual layered execution.
theorem execute_cells (dims : Nat → Nat) (N : Nat) :
(execute dims N).cellWrites = ∑ l ∈ Finset.range (N+1), (N+1-l) :=
buildLayers_cellWrites _ _ _Candidate visits come from the returned minimum-scan counters.
theorem execute_candidateVisits (dims : Nat → Nat) (N : Nat) :
(execute dims N).candidateVisits = ∑ l ∈ Finset.range (N+1), (N+1-l)*l :=
buildLayers_candidateVisits _ id _ (step_visits dims) _A finite stored table has the certified cost/split contract on its domain.
def CorrectTable (dims : Nat → Nat) (N : Nat) (rows : Array (Array Cell)) : Prop :=
∀ l i, l ≤ N → i+l ≤ N → CorrectCell dims l i (get rows l i)
private theorem stored_split {dims : Nat → Nat} {N : Nat} {rows : Array (Array Cell)}
(h : CorrectTable dims N rows) (i j : Nat) (hij : i < j) (hj : j ≤ N) :
i ≤ (get rows (j-i) i).split ∧ (get rows (j-i) i).split < j ∧
matrixChainOpt dims i j =
matrixSplitCost dims (matrixChainOpt dims) i j (get rows (j-i) i).split := by
have hc := h (j-i) i (by omega) (by omega)
have hs := hc.2 (by omega)
have he : i+(j-i) = j := by omega
dsimp only [CorrectCell] at hc
rw [he] at hc hs
exact ⟨hs.1, hs.2.1, hc.1.symm.trans hs.2.2⟩Reconstruct using stored indices only. The finite-domain proof is erased; no recursive optimum or legacy split selector is evaluated here.
def reconstruct (dims : Nat → Nat) (N : Nat) (rows : Array (Array Cell))
(h : CorrectTable dims N rows) (i j : Nat) (hij : i ≤ j) (hj : j ≤ N) : ChainPlan i j :=
if ht : i < j then
let k := (get rows (j-i) i).split
have hs := stored_split h i j ht hj
ChainPlan.split i k j
(reconstruct dims N rows h i k hs.1 (by omega))
(reconstruct dims N rows h (k+1) j (by omega) hj)
else
have he : i = j := by omega
he ▸ ChainPlan.single i
termination_by j-i
decreasing_by
all_goals
have _hs := stored_split h i j ht hj
try dsimp only [k] at *
omegaThe plan follows the stored split in every subinterval.
theorem reconstruct_reconstructed (dims : Nat → Nat) (N : Nat) (rows : Array (Array Cell))
(h : CorrectTable dims N rows) (i j : Nat) (hij : i ≤ j) (hj : j ≤ N) :
ChainPlan.ReconstructedBy (fun i j => (get rows (j-i) i).split)
(reconstruct dims N rows h i j hij hj) := by
rw [reconstruct]
split
next ht =>
have hs := stored_split h i j ht hj
exact ChainPlan.ReconstructedBy.split i _ j rfl
(reconstruct_reconstructed dims N rows h i _ hs.1 (by omega))
(reconstruct_reconstructed dims N rows h _ j (by omega) hj)
next ht =>
have he : i = j := by omega
subst j
exact ChainPlan.ReconstructedBy.single i
termination_by j-i
decreasing_by
all_goals
omegaStored-table reconstruction attains the independent optimum.
theorem reconstruct_cost (dims : Nat → Nat) (N : Nat) (rows : Array (Array Cell))
(h : CorrectTable dims N rows) (i j : Nat) (hij : i ≤ j) (hj : j ≤ N) :
ChainPlan.cost dims (reconstruct dims N rows h i j hij hj) = matrixChainOpt dims i j := by
rw [reconstruct]
split
next ht =>
have hs := stored_split h i j ht hj
rw [ChainPlan.cost, reconstruct_cost dims N rows h i _ hs.1 (by omega),
reconstruct_cost dims N rows h _ j (by omega) hj]
exact hs.2.2.symm
next ht =>
have he : i = j := by omega
subst j
simp [ChainPlan.cost, matrixChainOpt]
termination_by j-i
decreasing_by
all_goals
omegaBuild a table once and reconstruct the complete chain from that same table.
def executePlan (dims : Nat → Nat) (N : Nat) : ChainPlan 0 N × Table Cell :=
let table := execute dims N
(reconstruct dims N table.rows (fun l i hl hi => execute_get_correct dims N l i hl hi)
0 N (Nat.zero_le _) le_rfl, table)
theorem executePlan_correct (dims : Nat → Nat) (N : Nat) :
MatrixChainOptimalPlan dims (executePlan dims N).1 := by
intro other
have he := reconstruct_cost dims N (execute dims N).rows
(fun l i hl hi => execute_get_correct dims N l i hl hi) 0 N (Nat.zero_le _) le_rfl
change ChainPlan.cost dims (reconstruct dims N (execute dims N).rows _ 0 N _ _) ≤ _
rw [he]
exact matrixChain_opt_le_planCost (matrixChainOpt_lowerBound dims) otherEvery actual interval-evaluation event is unique.
theorem execute_once (dims : Nat → Nat) (N : Nat) :
(execute dims N).evaluatedStates.Nodup :=
buildLayers_evaluatedStates_nodup _ _ _The actual evaluation trace covers exactly the valid interval states.
theorem execute_state_iff (dims : Nat → Nat) (N l i : Nat) :
(l,i) ∈ (execute dims N).evaluatedStates ↔ i+l ≤ N := by
rw [execute, buildLayers_evaluatedStates_mem]
omegaQuadratic storage from actual interval writes.
theorem execute_cells_le (dims : Nat → Nat) (N : Nat) :
(execute dims N).cellWrites ≤ (N+1)^2 := by
rw [execute_cells]
exact interval_cells_le_square NCubic bound on the actual candidate-scan counter.
theorem execute_candidateVisits_le (dims : Nat → Nat) (N : Nat) :
(execute dims N).candidateVisits ≤ (N+1)^3 := by
rw [execute_candidateVisits]
exact interval_visits_le_cube NThe reconstructed plan's cost equals the optimum in the same returned table.
theorem executePlan_cost (dims : Nat → Nat) (N : Nat) :
ChainPlan.cost dims (executePlan dims N).1 =
(get (executePlan dims N).2.rows N 0).cost := by
dsimp only [executePlan]
rw [reconstruct_cost]
simpa using (execute_cost dims N 0 N (Nat.zero_le _) le_rfl).symmExact quadratic count of stored interval cells.
theorem execute_cells_closed (dims : Nat → Nat) (N : Nat) :
2 * (execute dims N).cellWrites = (N+1)*(N+2) := by
rw [execute_cells]
exact interval_cells_closed NExact cubic count of executed candidates.
theorem execute_candidateVisits_closed (dims : Nat → Nat) (N : Nat) :
6 * (execute dims N).candidateVisits = N*(N+1)*(N+2) := by
rw [execute_candidateVisits]
exact interval_visits_closed NTwo-sided quadratic storage bounds, including the diagonal boundary.
theorem execute_cells_bounds (dims : Nat → Nat) (N : Nat) :
(N+1)^2 ≤ 2 * (execute dims N).cellWrites ∧
(execute dims N).cellWrites ≤ (N+1)^2 := by
refine ⟨?_, execute_cells_le dims N⟩
rw [execute_cells_closed]
nlinarithTwo-sided cubic candidate bounds.
theorem execute_candidateVisits_bounds (dims : Nat → Nat) (N : Nat) :
N^3 ≤ 6 * (execute dims N).candidateVisits ∧
(execute dims N).candidateVisits ≤ (N+1)^3 := by
refine ⟨?_, execute_candidateVisits_le dims N⟩
rw [execute_candidateVisits_closed]
nlinarith [sq_nonneg (N : ℤ)]end CLRS.Chapter15.MatrixChainExecutionCLRSLean.Chapter_15.Section_15_2_Matrix_Chain_Multiplication
CLRS Section 15.2 - Matrix-chain multiplication
This section adds a first mathematical proof layer for matrix-chain
multiplication. A parenthesization is represented by an inductive
ChainPlan i j, and a candidate dynamic-programming cost table is
specified by the usual split lower bound. The main theorem says every concrete
parenthesization has cost at least the candidate optimum for its interval.
The file also adds a reconstruction certificate: if a split table records a
tight split for each nonsingleton interval, then any parenthesization rebuilt
from that split table has exactly the candidate optimal cost, and therefore has
cost no greater than any competing parenthesization. Any two plans
reconstructed from the same tight split table for the same interval have the
same cost.
Status: proved for the mathematical optimal-cost layer, recursive recurrence
evaluation, and optimal parenthesization.
Deferred refinements:
-
The recurrence evaluator repeats subintervals and is not a cached table. Fourth-edition Chapter 14 supplies separate interval-array executions and stored-selector reconstruction with attached cell/candidate counters.
namespace CLRSnamespace Chapter15Parenthesization model
A binary parenthesization of the matrix chain from index i to j.
inductive ChainPlan : Nat → Nat → Type where
| single (i : Nat) : ChainPlan i i
| split (i k j : Nat) :
ChainPlan i k → ChainPlan (k + 1) j → ChainPlan i jnamespace ChainPlanThe left endpoint of a chain plan is at most the right endpoint.
theorem start_le_end {i j : Nat} (plan : ChainPlan i j) : i ≤ j := by
induction plan with
| single i =>
exact le_rfl
| split i k j left right ihLeft ihRight =>
exact Nat.le_trans ihLeft (Nat.le_trans (Nat.le_succ k) ihRight)
The scalar multiplication cost of a parenthesization, using the CLRS dimension
array convention: matrix A_i has dimensions dims i by
dims (i+1).
def cost (dims : Nat → Nat) : {i j : Nat} → ChainPlan i j → Nat
| _, _, single _ => 0
| _, _, split i k j left right =>
cost dims left + cost dims right + dims i * dims (k + 1) * dims (j + 1)A parenthesization is reconstructed from a split table when every internal node uses the split index prescribed for its interval.
inductive ReconstructedBy (splitAt : Nat → Nat → Nat) :
{i j : Nat} → ChainPlan i j → Prop where
| single (i : Nat) : ReconstructedBy splitAt (single i)
| split (i k j : Nat) {left : ChainPlan i k}
{right : ChainPlan (k + 1) j} :
k = splitAt i j →
ReconstructedBy splitAt left →
ReconstructedBy splitAt right →
ReconstructedBy splitAt (ChainPlan.split i k j left right)end ChainPlanOptimality interface
The CLRS split cost for multiplying matrices i..j split after k.
def matrixSplitCost (dims : Nat → Nat) (opt : Nat → Nat → Nat)
(i j k : Nat) : Nat :=
opt i k + opt (k + 1) j + dims i * dims (k + 1) * dims (j + 1)A candidate cost table satisfies the matrix-chain lower-bound recurrence if every valid first split has cost at least the table entry.
def MatrixChainLowerBound (dims : Nat → Nat) (opt : Nat → Nat → Nat) : Prop :=
(∀ i, opt i i = 0) ∧
∀ {i j k}, k ∈ Finset.Icc i (j - 1) →
opt i j ≤ matrixSplitCost dims opt i j kA split table is tight for a candidate matrix-chain cost table when each nonsingleton interval chooses a valid first split whose split cost is exactly the table entry.
def MatrixChainSplitOptimal (dims : Nat → Nat) (opt : Nat → Nat → Nat)
(splitAt : Nat → Nat → Nat) : Prop :=
(∀ i, opt i i = 0) ∧
∀ {i j}, i < j →
splitAt i j ∈ Finset.Icc i (j - 1) ∧
opt i j = matrixSplitCost dims opt i j (splitAt i j)A concrete parenthesization is optimal when no other plan has lower cost.
def MatrixChainOptimalPlan (dims : Nat → Nat)
{i j : Nat} (plan : ChainPlan i j) : Prop :=
∀ other : ChainPlan i j, ChainPlan.cost dims plan ≤ ChainPlan.cost dims otherEvery concrete parenthesization has cost at least the candidate optimum specified by the recurrence lower-bound interface.
theorem matrixChain_opt_le_planCost {dims : Nat → Nat}
{opt : Nat → Nat → Nat} (hopt : MatrixChainLowerBound dims opt) :
∀ {i j : Nat} (plan : ChainPlan i j),
opt i j ≤ ChainPlan.cost dims plan := by
intro i j plan
induction plan with
| single i =>
simpa [ChainPlan.cost] using hopt.1 i
| split i k j left right ihLeft ihRight =>
have hik : i ≤ k := ChainPlan.start_le_end left
have hkj : k + 1 ≤ j := ChainPlan.start_le_end right
have hmem : k ∈ Finset.Icc i (j - 1) := by
rw [Finset.mem_Icc]
omega
have hsplit := hopt.2 hmem
unfold matrixSplitCost at hsplit
simp [ChainPlan.cost]
omegaAny plan reconstructed from a tight split table has exactly the candidate optimal cost.
theorem matrixChain_reconstructed_cost_eq {dims : Nat → Nat}
{opt : Nat → Nat → Nat} {splitAt : Nat → Nat → Nat}
(hsplit : MatrixChainSplitOptimal dims opt splitAt) :
∀ {i j : Nat} {plan : ChainPlan i j},
ChainPlan.ReconstructedBy splitAt plan →
ChainPlan.cost dims plan = opt i j := by
intro i j plan hrec
induction hrec with
| single i =>
simpa [ChainPlan.cost] using (hsplit.1 i).symm
| split =>
rename_i i k j left right hk _hleft _hright ihLeft ihRight
subst k
have hij : i < j := by
have hleftLe : i ≤ splitAt i j := ChainPlan.start_le_end left
have hrightLe : splitAt i j + 1 ≤ j := ChainPlan.start_le_end right
omega
rcases hsplit.2 hij with ⟨_hmem, hcost⟩
simp [ChainPlan.cost, matrixSplitCost, ihLeft, ihRight, hcost]Any two parenthesizations reconstructed from the same tight split table for the same interval have equal cost.
theorem matrixChain_reconstructed_cost_eq_of_reconstructed {dims : Nat → Nat}
{opt : Nat → Nat → Nat} {splitAt : Nat → Nat → Nat}
(hsplit : MatrixChainSplitOptimal dims opt splitAt)
{i j : Nat} {left right : ChainPlan i j}
(hleft : ChainPlan.ReconstructedBy splitAt left)
(hright : ChainPlan.ReconstructedBy splitAt right) :
ChainPlan.cost dims left = ChainPlan.cost dims right := by
calc
ChainPlan.cost dims left = opt i j :=
matrixChain_reconstructed_cost_eq hsplit hleft
_ = ChainPlan.cost dims right :=
(matrixChain_reconstructed_cost_eq hsplit hright).symmCombining a lower-bound table with a tight split-table reconstruction proves the reconstructed parenthesization is globally optimal.
theorem matrixChain_reconstructed_optimal {dims : Nat → Nat}
{opt : Nat → Nat → Nat} {splitAt : Nat → Nat → Nat}
(hlower : MatrixChainLowerBound dims opt)
(hsplit : MatrixChainSplitOptimal dims opt splitAt)
{i j : Nat} {plan : ChainPlan i j}
(hrec : ChainPlan.ReconstructedBy splitAt plan) :
MatrixChainOptimalPlan dims plan := by
intro other
have hcost :
ChainPlan.cost dims plan = opt i j :=
matrixChain_reconstructed_cost_eq hsplit hrec
have hother :
opt i j ≤ ChainPlan.cost dims other :=
matrixChain_opt_le_planCost hlower other
omegaDirect cost inequality form of the split-table reconstruction theorem: a plan rebuilt from a tight split table is no more expensive than any other parenthesization of the same interval.
theorem matrixChain_reconstructed_cost_le_planCost {dims : Nat → Nat}
{opt : Nat → Nat → Nat} {splitAt : Nat → Nat → Nat}
(hlower : MatrixChainLowerBound dims opt)
(hsplit : MatrixChainSplitOptimal dims opt splitAt)
{i j : Nat} {plan : ChainPlan i j}
(hrec : ChainPlan.ReconstructedBy splitAt plan)
(other : ChainPlan i j) :
ChainPlan.cost dims plan ≤ ChainPlan.cost dims other := by
exact matrixChain_reconstructed_optimal hlower hsplit hrec otherRecursive cost specification and final optimality
def matrixChainOpt (dims : Nat → Nat) : Nat → Nat → Nat
| i, j =>
if h : i < j then
(Finset.Icc i (j - 1)).attach.inf'
(Finset.attach_nonempty_iff.mpr (by
use i; simp [Finset.mem_Icc]; omega))
(fun k =>
matrixChainOpt dims i k.1 +
matrixChainOpt dims (k.1 + 1) j +
dims i * dims (k.1 + 1) * dims (j + 1))
else
0
termination_by i j => j - i
decreasing_by
all_goals
have hk := Finset.mem_Icc.mp k.2
omega
theorem matrixChainOpt_lowerBound (dims : Nat → Nat) :
MatrixChainLowerBound dims (matrixChainOpt dims) := by
refine ⟨?_, ?_⟩
· intro i; unfold matrixChainOpt; simp
· intro i j k hk
rcases Finset.mem_Icc.mp hk with ⟨hik, hkj⟩
by_cases hij : i < j
· unfold matrixChainOpt; simp [hij]
unfold matrixSplitCost
have hm : (⟨k, Finset.mem_Icc.mpr ⟨hik, hkj⟩⟩ : {x // x ∈ Finset.Icc i (j - 1)}) ∈
(Finset.Icc i (j - 1)).attach := by simp
-- Goal: attach.inf' (fun r => ... r.1 ...) ≤ matrixChainOpt i k + ...
-- Finset.inf'_le gives: attach.inf' (fun r => ...) ≤ (fun r => ... r.1 ...) ⟨k, ...⟩
-- After beta reduction, RHS = matrixChainOpt i k + ...
let f : {x // x ∈ Finset.Icc i (j - 1)} → ℕ :=
λ r => matrixChainOpt dims i r.1 + matrixChainOpt dims (r.1 + 1) j +
dims i * dims (r.1 + 1) * dims (j + 1)
simpa [f] using Finset.inf'_le f hm
· have hzero : matrixChainOpt dims i j = 0 := by
unfold matrixChainOpt; simp [hij]
rw [hzero]
exact Nat.zero_le _
private lemma bridge_attach_inf (dims : Nat → Nat) (i j : Nat) (hij : i < j) :
matrixChainOpt dims i j =
(Finset.Icc i (j - 1)).inf'
(by use i; simp [Finset.mem_Icc]; omega)
(λ k => matrixChainOpt dims i k + matrixChainOpt dims (k + 1) j +
dims i * dims (k + 1) * dims (j + 1)) := by
let s := Finset.Icc i (j - 1)
have Hs : s.Nonempty := by use i; simp [s, Finset.mem_Icc]; omega
let f (k : ℕ) := matrixChainOpt dims i k + matrixChainOpt dims (k + 1) j +
dims i * dims (k + 1) * dims (j + 1)
let g (r : {x // x ∈ s}) : ℕ :=
matrixChainOpt dims i r.1 + matrixChainOpt dims (r.1 + 1) j +
dims i * dims (r.1 + 1) * dims (j + 1)
have h1 : matrixChainOpt dims i j = s.attach.inf' (Finset.attach_nonempty_iff.mpr Hs) g := by
unfold matrixChainOpt; simp [hij, s, g]
have h2 : s.attach.inf' (Finset.attach_nonempty_iff.mpr Hs) g = s.inf' Hs f := by
have h_att : s.attach.Nonempty := Finset.attach_nonempty_iff.mpr Hs
apply le_antisymm
· -- attach.inf' ≤ inf', via lower bound on inf'
apply Finset.le_inf' Hs f
intro x hx
have hm : (⟨x, hx⟩ : {x // x ∈ s}) ∈ s.attach := by simp
simpa [f, g] using Finset.inf'_le g hm
· -- inf' ≤ attach.inf', via lower bound on attach.inf'
apply Finset.le_inf' h_att g
intro r hr
-- r ∈ s.attach, so r.2 : r.1 ∈ s
simpa [f, g] using Finset.inf'_le f r.2
rw [h1, h2]
private lemma exists_inf'_eq (s : Finset ℕ) (h : s.Nonempty) (f : ℕ → ℕ) :
∃ a ∈ s, f a = s.inf' h f := by
induction' s using Finset.induction with a s has ih
· exact absurd h (by simp)
· by_cases hs : s.Nonempty
· rcases ih hs with ⟨b, hb, hb_eq⟩
rw [Finset.inf'_insert hs f]
by_cases hle : f a ≤ s.inf' hs f
· rw [min_eq_left hle]
exact ⟨a, Finset.mem_insert_self a s, rfl⟩
· rw [min_eq_right (by omega : s.inf' hs f ≤ f a)]
rw [← hb_eq]
exact ⟨b, Finset.mem_insert_of_mem hb, rfl⟩
· have hsingleton : s = ∅ := Finset.not_nonempty_iff_eq_empty.mp hs
subst hsingleton
simp
The computable split-point selector for matrix-chain DP. For interval
i < j, it selects the smallest k in [i, j-1] that attains
the minimum split cost. This makes the entire reconstruction chain computable
without Exists.choose.
def matrixChainSplit (dims : Nat → Nat) (i j : Nat) : Nat :=
if h : i < j then
(Finset.Icc i (j - 1)).filter
(fun k =>
matrixChainOpt dims i k + matrixChainOpt dims (k + 1) j +
dims i * dims (k + 1) * dims (j + 1) = matrixChainOpt dims i j) |>.min'
(by
have h_nonempty : (Finset.Icc i (j - 1)).Nonempty := by
use i; simp [Finset.mem_Icc]; omega
let f (k : ℕ) := matrixChainOpt dims i k + matrixChainOpt dims (k + 1) j +
dims i * dims (k + 1) * dims (j + 1)
have h_eq : matrixChainOpt dims i j =
(Finset.Icc i (j - 1)).inf' h_nonempty f :=
bridge_attach_inf dims i j h
have h_exists := exists_inf'_eq (Finset.Icc i (j - 1)) h_nonempty f
rw [← h_eq] at h_exists
rcases h_exists with ⟨k, hk, hk_eq⟩
exact ⟨k, Finset.mem_filter.mpr ⟨hk, hk_eq⟩⟩)
else
i
theorem matrixChainSplit_optimal (dims : Nat → Nat) (i j : Nat) (hij : i < j) :
matrixChainSplit dims i j ∈ Finset.Icc i (j - 1) ∧
matrixChainOpt dims i j =
matrixSplitCost dims (matrixChainOpt dims) i j (matrixChainSplit dims i j) := by
let s : Finset ℕ := Finset.Icc i (j - 1)
have h_nonempty : s.Nonempty := by use i; simp [s, Finset.mem_Icc]; omega
let f (k : ℕ) := matrixChainOpt dims i k + matrixChainOpt dims (k + 1) j +
dims i * dims (k + 1) * dims (j + 1)
have h_opt_eq : matrixChainOpt dims i j = s.inf' h_nonempty f :=
bridge_attach_inf dims i j hij
have h_exists : ∃ k ∈ s, f k = matrixChainOpt dims i j := by
rw [h_opt_eq]; exact exists_inf'_eq s h_nonempty f
have h_filter_nonempty : (s.filter fun k => f k = matrixChainOpt dims i j).Nonempty := by
rcases h_exists with ⟨k, hk, hk_eq⟩
exact ⟨k, Finset.mem_filter.mpr ⟨hk, hk_eq⟩⟩
set k := (s.filter fun k => f k = matrixChainOpt dims i j).min' h_filter_nonempty with hk_def
have hk_mem_filter : k ∈ s.filter fun k => f k = matrixChainOpt dims i j := by
rw [hk_def]; exact Finset.min'_mem _ h_filter_nonempty
have hk_mem : k ∈ s := (Finset.mem_filter.mp hk_mem_filter).1
have hk_eq : f k = matrixChainOpt dims i j := (Finset.mem_filter.mp hk_mem_filter).2
have h_split_val : matrixChainSplit dims i j = k := by
unfold matrixChainSplit
simp [hij, s, f, hk_def]
rw [h_split_val]
refine ⟨hk_mem, ?_⟩
unfold matrixSplitCost
dsimp [f] at hk_eq
rw [← hk_eq]theorem matrixChainOpt_splitOptimal (dims : Nat → Nat) :
MatrixChainSplitOptimal dims (matrixChainOpt dims) (matrixChainSplit dims) := by
refine ⟨?_, ?_⟩
· intro i; simp [matrixChainOpt]
· intro i j hij; exact matrixChainSplit_optimal dims i j hijdef matrixChainReconstruct (dims : Nat → Nat) (i j : Nat) (hbound : i ≤ j) : ChainPlan i j :=
if h : i < j then
let k := matrixChainSplit dims i j
ChainPlan.split i k j
(matrixChainReconstruct dims i k (by
have hsplit := matrixChainSplit_optimal dims i j h
rcases hsplit with ⟨hmem, _⟩
rcases Finset.mem_Icc.mp hmem with ⟨hlo, hhi⟩
omega))
(matrixChainReconstruct dims (k + 1) j (by
have hsplit := matrixChainSplit_optimal dims i j h
rcases hsplit with ⟨hmem, _⟩
rcases Finset.mem_Icc.mp hmem with ⟨hlo, hhi⟩
omega))
else
have heq : i = j := by omega
heq ▸ ChainPlan.single j
termination_by j - i
decreasing_by
· have hsplit := matrixChainSplit_optimal dims i j h
rcases hsplit with ⟨hmem, _⟩
rcases Finset.mem_Icc.mp hmem with ⟨hlo, _hhi⟩
omega
· have hsplit := matrixChainSplit_optimal dims i j h
rcases hsplit with ⟨hmem, _⟩
rcases Finset.mem_Icc.mp hmem with ⟨hlo, _hhi⟩
omega
theorem matrixChainReconstruct_reconstructed (dims : Nat → Nat) (i j : Nat) (hbound : i ≤ j) :
ChainPlan.ReconstructedBy (matrixChainSplit dims) (matrixChainReconstruct dims i j hbound) := by
unfold matrixChainReconstruct
split
· next h =>
have hsplit := matrixChainSplit_optimal dims i j h
rcases hsplit with ⟨hmem, _⟩
rcases Finset.mem_Icc.mp hmem with ⟨hlo, hhi⟩
have h_left_bound : i ≤ matrixChainSplit dims i j := by omega
have h_right_bound : matrixChainSplit dims i j + 1 ≤ j := by omega
simp
refine ChainPlan.ReconstructedBy.split i (matrixChainSplit dims i j) j rfl
(matrixChainReconstruct_reconstructed dims i (matrixChainSplit dims i j) (by omega))
(matrixChainReconstruct_reconstructed dims (matrixChainSplit dims i j + 1) j (by omega))
· next h =>
have heq : i = j := by omega
cases heq
exact ChainPlan.ReconstructedBy.single i
termination_by j - i
decreasing_by
· omega
· omegatheorem matrixChain_correct (dims : Nat → Nat) (i j : Nat) (hbound : i ≤ j) :
∃ plan : ChainPlan i j, MatrixChainOptimalPlan dims plan := by
let plan := matrixChainReconstruct dims i j hbound
refine ⟨plan, matrixChain_reconstructed_optimal
(matrixChainOpt_lowerBound dims)
(matrixChainOpt_splitOptimal dims)
(matrixChainReconstruct_reconstructed dims i j hbound)⟩end Chapter15end CLRS