Repository navigation
Expand file tree
/
Copy pathMatN.v
More file actions
2124 lines (1875 loc) · 78.7 KB
/
Copy pathMatN.v
File metadata and controls
2124 lines (1875 loc) · 78.7 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
789
790
791
792
793
794
795
796
797
798
799
800
801
802
803
804
805
806
807
808
809
810
811
812
813
814
815
816
817
818
819
820
821
822
823
824
825
826
827
828
829
830
831
832
833
834
835
836
837
838
839
840
841
842
843
844
845
846
847
848
849
850
851
852
853
854
855
856
857
858
859
860
861
862
863
864
865
866
867
868
869
870
871
872
873
874
875
876
877
878
879
880
881
882
883
884
885
886
887
888
889
890
891
892
893
894
895
896
897
898
899
900
901
902
903
904
905
906
907
908
909
910
911
912
913
914
915
916
917
918
919
920
921
922
923
924
925
926
927
928
929
930
931
932
933
934
935
936
937
938
939
940
941
942
943
944
945
946
947
948
949
950
951
952
953
954
955
956
957
958
959
960
961
962
963
964
965
966
967
968
969
970
971
972
973
974
975
976
977
978
979
980
981
982
983
984
985
986
987
988
989
990
991
992
993
994
995
996
997
998
999
1000
(** Matrix Operations over a Finite Semiring (list-based)
File: MatN.v
Matrix is `Node -> Node -> R` where `Node` is a finite type
and `R` is a semiring. This file provides both functional
(high‑level) and list‑based (computationally efficient) matrix
operations together with proofs of their equivalence. *)
From Stdlib Require Import List Utf8
BinNatDef
Lia PeanoNat PArith.
From Semiring Require Import OrelN Structures
PathN.
Import ListNotations SemiringNotations.
(** Section: Generic Definitions
General‑purpose combinators on lists used by the matrix
operations below. These are polymorphic in the element type. *)
Section GenericDef.
(** [zip_with f xs ys] applies `f` pointwise to the elements of `xs`
and `ys`, stopping at the shorter list. This is the standard
"zipWith" from functional programming. *)
(** Element-wise combination of two lists; stops at the shorter list. *)
Fixpoint zip_with {A B C : Type}
(f : A -> B -> C) (xs : list A) (ys : list B) : list C :=
match xs, ys with
| x :: xs, y :: ys => f x y :: zip_with f xs ys
| _, _ => []
end.
(** [transpose_list xss] transposes a rectangular list of lists.
For a singleton row [[r]] it returns the column list
[[x₁];[x₂];…]; for multiple rows it uses [zip_with cons]
to build the transposed rows incrementally. *)
(** Transpose a rectangular list-of-lists (matrix). *)
Fixpoint transpose_list {A : Type} (xss : list (list A)) : list (list A) :=
match xss with
| [] => []
| xssh :: xsst =>
match xsst with
| [] => map (fun y => [y]) xssh
| _ :: _ => zip_with List.cons xssh (transpose_list xsst)
end
end.
End GenericDef.
(** * Section: Proofs about the generic list combinators *)
Section GenericDefProofs.
(** The length of [zip_with f xs ys] is the minimum of the lengths
of [xs] and [ys]. *)
(** The length of [zip_with] is the minimum of the two input lengths. *)
Theorem zip_with_length {A B C : Type}
(f : A -> B -> C) : ∀ (xs : list A) (ys : list B),
List.length (zip_with f xs ys) =
Nat.min (List.length xs) (List.length ys).
Proof.
induction xs as [|xsh xst ih].
+
intros ys; reflexivity.
+
intros *.
destruct ys as [|ysh yst].
++
cbn; reflexivity.
++
cbn. f_equal.
eapply ih.
Qed.
(** If three lists have equal length, [zip_with cons] preserves that length. *)
Theorem zip_transpose_length {A : Type} :
∀ (xs ys : list A) zs,
List.length xs = List.length ys ->
List.length ys = List.length zs ->
length xs = length (zip_with cons ys zs).
Proof.
induction xs as [|xsh xst ih].
+
intros [|ysh yst] [|zsh zst] ha hb;
cbn in ha, hb; try congruence;
try reflexivity.
+
intros [|ysh yst] [|zsh zst] ha hb;
cbn in ha, hb; try congruence.
cbn. erewrite <-ih.
reflexivity.
inversion ha;
reflexivity.
inversion hb; reflexivity.
Qed.
(** In a rectangular matrix, row length equals the length of the transpose. *)
Theorem transpose_length {A : Type} :
∀ (xst : list (list A)) (xsh : list A),
0 < List.length xst -> 0 < List.length xsh ->
(∀ xs : list A, In xs (xsh :: xst) → ∀ ys : list A,
In ys (xsh :: xst) → length xs = length ys ∧ 0 < length xs) ->
(** transpose_eff (transpose_eff xst) = xst -> *)
length xsh = length (transpose_list xst).
Proof.
induction xst as [|xsth xstt ih].
+
intros * ha hb hc.
cbn in ha; nia.
+
destruct xstt as [|xsstth xssttt].
++
intros [|xshh xsht] ha hb hc.
*
cbn in hb; nia.
*
pose proof (hc (xshh :: xsht) (or_introl eq_refl)
xsth (or_intror (or_introl eq_refl))) as he.
destruct he as (hel & her).
cbn. rewrite length_map, <-hel;
reflexivity.
++
(** induction case *)
assert (hd : transpose_list (xsth :: xsstth :: xssttt) =
zip_with List.cons xsth (transpose_list (xsstth :: xssttt))).
cbn. reflexivity.
remember (xsstth :: xssttt) as xst.
intros * ha hb hc.
rewrite hd.
destruct (hc xsh (or_introl eq_refl) xsth
(or_intror (or_introl eq_refl))) as (hel & her).
rewrite Heqxst in hc.
destruct (hc xsth (or_intror (or_introl eq_refl))
xsstth (or_intror (or_intror (or_introl eq_refl)))) as
(hfl & hfr).
assert (hg : 0 < length xst). subst; cbn; nia.
assert(hf : (∀ xs : list A, In xs (xsh :: xst) → ∀ ys : list A,
In ys (xsh :: xst) → length xs = length ys ∧ 0 < length xs)).
{
intros * hf * hi.
apply hc.
rewrite <- Heqxst.
firstorder.
rewrite <-Heqxst.
firstorder.
}
specialize (ih xsh hg hb hf).
rewrite Heqxst.
eapply zip_transpose_length.
assumption.
rewrite <-Heqxst.
nia.
Qed.
(** Helper: nth does not depend on default when index is in bounds. *)
Lemma nth_default_indep :
forall (A : Type) (idx : nat) (l : list A) (d1 d2 : A),
(idx < List.length l)%nat -> List.nth idx l d1 = List.nth idx l d2.
Proof.
intros A idx l d1 d2 Hlt.
generalize dependent idx.
induction l as [|a l' IH]; intros idx Hlt.
- inversion Hlt.
- destruct idx as [|idx'].
+ reflexivity.
+ simpl in Hlt. cbn.
eapply IH. nia.
Qed.
(** Helper lemma: nth 0 of nth i on map singletons = nth i on original *)
Lemma nth_0_map_singleton :
forall (A : Type) (l : list A) (i : nat) (d : A),
List.nth 0 (List.nth i (List.map (fun y => [y]) l) ([] : list A)) d =
List.nth i l d.
Proof.
intros A l i d. revert i.
induction l as [|x l' IHl]; intros i; simpl.
- destruct i; reflexivity.
- destruct i as [|i']; simpl; [reflexivity | apply IHl].
Qed.
End GenericDefProofs.
Section Matrix.
Context
{Node : FinType.type}.
(** A matrix over semiring [R] indexed by finite type [Node]. *)
Let Matrix {R : Semiring.type} := @OrelN.Matrix Node R.
(** returns the cth row of m *)
Definition row {R : Semiring.type} (m : Matrix) (c : Node) : Node -> R :=
fun d => m c d.
(** returns the cth column of m *)
Definition col {R : Semiring.type} (m : Matrix) (c : Node) : Node -> R :=
fun d => m d c.
(** zero matrix, additive identity of plus *)
Definition zeroM {R : Semiring.type} : @Matrix R :=
fun _ _ => 0.
(** identity matrix, mulitplicative identity of mul
Idenitity Matrix *)
Definition I {R : Semiring.type} : @Matrix R :=
fun (c d : Node) =>
match fin_eq_dec c d with
| left _ => 1
| right _ => 0
end.
(** transpose the matrix m *)
Definition transpose {R : Semiring.type} (m : @Matrix R) : @Matrix R :=
fun (c d : Node) => m d c.
(** pointwise addition to two matrices *)
Definition matrix_add {R : Semiring.type} (m₁ m₂ : @Matrix R) : @Matrix R :=
fun c d => (m₁ c d + m₂ c d).
(** Finite sum of a [Node]-indexed family over the semiring. *)
Definition sum {R : Semiring.type} (f : Node -> R) : R :=
List.fold_right (fun x y => f x + y) 0 elements.
(** Extensionality of [sum]: equal functions have equal sums. *)
Lemma sum_ext {R : Semiring.type} : forall (f g : Node -> R),
(forall x, f x = g x) -> sum f = sum g.
Proof.
intros f g Heq.
unfold sum.
induction elements as [|a l IH]; simpl.
- reflexivity.
- rewrite Heq. f_equal. exact IH.
Qed.
(** generalised matrix multiplication *)
Definition matrix_mul {R : Semiring.type}
(m₁ m₂ : @Matrix R) : @Matrix R:=
fun (c d : Node) =>
sum (fun y => (m₁ c y * m₂ y d)).
(** ** Matrix exponentiation *)
Local Infix "+M" := matrix_add (at level 50, only parsing).
(** Linear matrix exponentiation: [pow m n = m * m * ... * m] (n times). *)
Fixpoint pow {R : Semiring.type} (m : @Matrix R) (n : nat) : @Matrix R :=
match n with
| 0%nat => I
| S n' => matrix_mul m (pow m n')
end.
(** Linear matrix exponentiation: [pow m n = m * m * ... * m] (n times). *)
Fixpoint pow_pos {R : Semiring.type} (e : @Matrix R) (n : positive) : @Matrix R :=
match n with
| xH => e
| xO p => let ret := pow_pos e p in matrix_mul ret ret
| xI p =>
let reta := pow_pos e p in
let retb := matrix_mul reta reta in
matrix_mul e retb
end.
(** Matrix exponentiation for [N] (binary for positive, identity for zero). *)
Definition powN {R : Semiring.type} (e : @Matrix R) (n : N) : @Matrix R :=
match n with
| N0 => I
| Npos p => pow_pos e p
end.
(** ** Scalar exponentiation and partial sums *)
Fixpoint scalar_pow {R : Semiring.type} (a : R) (n : nat) : R :=
match n with
| O => 1
| S n' => a * scalar_pow a n'
end.
(** Scalar geometric series: [1 + a + a² + ... + aⁿ]. *)
Fixpoint scalar_geom_sum {R : Semiring.type} (a : R) (n : nat) : R :=
match n with
| O => 1
| S n' => (scalar_geom_sum a n') + scalar_pow a n
end.
(** Matrix geometric series: [I + M + M² + ... + Mⁿ]. *)
Fixpoint geom_sum {R : Semiring.type} (m : @Matrix R) (n : nat) : @Matrix R :=
match n with
| O => I
| S n' => (geom_sum m n') +M (pow m n)
end.
(** ** Efficient list-based matrix operations *)
(** Dot product of two lists *)
Definition dot_product {R : Semiring.type} (v1 v2 : list R) : R :=
fold_left add (map (fun '(x, y) => mul x y)
(combine v1 v2)) zero.
(** Matrix multiplication (list-based) *)
Definition mul_list {R : Semiring.type} (la lb : list (list R)) : list (list R) :=
let lbT := transpose_list lb in
map (fun row =>
map (fun col => dot_product row col) lbT) la.
(** Linear matrix exponentiation: [pow m n = m * m * ... * m] (n times). *)
Fixpoint pow_list {R : Semiring.type} (m : list (list R))
(n : nat) : list (list R) :=
match n with
| 0%nat => List.map (fun r => List.map (fun c => I r c) elements) elements
| S n' => mul_list m (pow_list m n')
end.
(** Linear matrix exponentiation: [pow m n = m * m * ... * m] (n times). *)
Fixpoint pow_pos_list {R : Semiring.type} (e : list (list R))
(n : positive) : list (list R) :=
match n with
| xH => e
| xO p => let ret := pow_pos_list e p in
mul_list ret ret
| xI p => let reta := pow_pos_list e p in
let retb := mul_list reta reta in
mul_list e retb
end.
(** Matrix exponentiation for [N] (binary for positive, identity for zero). *)
Definition powN_list {R : Semiring.type} (e : list (list R)) (n : N) : list (list R) :=
match n with
| N0 => List.map (fun r => List.map (fun c => I r c) elements) elements
| Npos p => pow_pos_list e p
end.
(** ** Lookup & conversion helpers (mirrors list_lookup in Semimodule) *)
(** Boolean decidable equality on Node *)
Definition eq_decb (x y : Node) : bool :=
match fin_eq_dec x y with left _ => true | right _ => false end.
(** Parallel list lookup keyed by finN with a default value *)
Fixpoint list_lookup {A : Type} (def : A)
(keys : list Node) (vals : list A) (key : Node) : A :=
match keys, vals with
| k :: ks, v :: vs => if eq_decb key k then v else list_lookup def ks vs key
| _, _ => def
end.
(** Convert between functional and list-of-lists representations *)
Definition to_list {R : Semiring.type} (m : Node -> Node -> R) : list (list R) :=
List.map (fun r => List.map (fun c => m r c) elements) elements.
(** Reconstruct a functional matrix from a list-of-lists representation. *)
Definition of_list {R : Semiring.type} (me : list (list R)) : @Matrix R :=
fun r c =>
let row := list_lookup [] elements me r in
list_lookup 0 elements row c.
(** Functional matrix exponentiation via the list-based implementation. *)
Definition pow_fun {R : Semiring.type} (m : Node -> Node -> R) (n : nat)
: Node -> Node -> R :=
of_list (pow_list (to_list m) n).
(** Matrix exponentiation for [N] (binary for positive, identity for zero). *)
Definition powN_fun {R : Semiring.type} (m : Node -> Node -> R) (n : N)
: Node -> Node -> R :=
of_list (powN_list (to_list m) n).
(** ** Correctness: looking up a tabulated matrix returns the original *)
Lemma list_lookup_map_aux : forall (A : Type) (def : A) (f : Node -> A)
(y : Node) (l : list Node), NoDup l -> In y l ->
list_lookup def l (List.map f l) y = f y.
Proof.
intros A def f y l Hnd Hin.
revert y Hnd Hin.
induction l as [|k ks IH]; intros y Hnd Hin.
- inversion Hin.
- inversion Hnd as [|nd_x nd_l Hnin Hnd_ks]; subst.
simpl in Hin. destruct Hin as [Hx | Hin_ks].
+ subst y.
unfold list_lookup, eq_decb.
destruct (fin_eq_dec k k) as [Heq | Hneq].
* reflexivity.
* exfalso. apply Hneq; reflexivity.
+ unfold list_lookup, eq_decb.
destruct (fin_eq_dec y k) as [Heq | Hneq].
* subst y. exfalso. exact (Hnin Hin_ks).
* apply IH; [exact Hnd_ks | exact Hin_ks].
Qed.
(** Looking up in a tabulated list returns the function value. *)
Lemma list_lookup_map : forall (A : Type) (def : A) (f : Node -> A)
(x : Node), list_lookup def elements (List.map f elements) x = f x.
Proof.
intros A def f x.
apply list_lookup_map_aux.
- apply (elements_nodup (s := Node)).
- apply elements_complete.
Qed.
(** Round-trip: [of_list (to_list m) = m]. *)
Lemma of_list_to_list {R : Semiring.type} : forall (m : @Matrix R) (r c : Node),
of_list (to_list m) r c = m r c.
Proof.
intros m r c.
unfold of_list, to_list.
simpl.
(** list_lookup 0 elements (list_lookup [] elements
(map (fun r0 => map (fun c0 => m r0 c0) elements) elements) r) c = m r c *)
rewrite list_lookup_map with (def := [])
(f := fun r => List.map (fun c => m r c) elements).
(** Now: list_lookup 0 elements (map (fun c0 => m r c0) elements) c = m r c *)
apply list_lookup_map.
Qed.
(** ** pow = pow_fun (mathematical = efficient, unary case) *)
(** ** Helper: positional index of a node in elements *)
Fixpoint index_of_aux (k : Node) (xs : list Node) : nat :=
match xs with
| [] => 0
| x :: xs' => if fin_eq_dec k x then 0 else S (index_of_aux k xs')
end.
(** Positional index of a node in the canonical element list. *)
Definition index_of (k : Node) : nat :=
index_of_aux k (elements (s := Node)).
(** The index of any node is within the element list length. *)
Lemma index_of_bound : forall (k : Node),
(index_of k < List.length (elements (s := Node)))%nat.
Proof.
intros k. unfold index_of.
assert (Hin : In k (elements (s := Node))) by apply elements_complete.
revert k Hin.
induction (elements (s := Node)) as [|e es IH]; intros k Hin; simpl.
- inversion Hin.
- destruct (fin_eq_dec k e) as [Heq | Hneq].
+ subst; simpl; nia.
+ simpl in Hin. destruct Hin as [Heq' | Hin']; [exfalso; apply Hneq; symmetry; exact Heq' |].
simpl. specialize (IH k Hin'). nia.
Qed.
(** ** list_lookup expressed via nth + index_of *)
Lemma list_lookup_nth_gen_aux : forall (A : Type) (d : A) (xs : list Node) (l : list A) (k : Node),
NoDup xs -> In k xs ->
list_lookup d xs l k = List.nth (index_of_aux k xs) l d.
Proof.
induction xs as [|e es IH]; intros l k Hnd Hin.
- inversion Hin.
- simpl in Hin. destruct Hin as [Hk | Hin'].
+ (* k = e *)
subst k.
destruct l as [|a l']; simpl.
* destruct (fin_eq_dec e e); reflexivity.
* unfold list_lookup, eq_decb. simpl.
destruct (fin_eq_dec e e) as [_ | Hneq]; [reflexivity | exfalso; apply Hneq; reflexivity].
+ (* k ∈ es *)
inversion Hnd as [|? ? Hnin Hnd']; subst.
destruct l as [|a l']; simpl.
* destruct (fin_eq_dec k e); reflexivity.
* unfold list_lookup, eq_decb. simpl.
destruct (fin_eq_dec k e) as [Heq | Hneq].
-- subst k. exfalso. apply Hnin. exact Hin'.
-- simpl. apply (IH l' k Hnd' Hin').
Qed.
(** [list_lookup] is equivalent to [nth] at the index position. *)
Lemma list_lookup_nth_gen : forall (A : Type) (d : A) (l : list A) (k : Node),
list_lookup d elements l k = List.nth (index_of k) l d.
Proof.
intros A d l k.
unfold index_of.
apply list_lookup_nth_gen_aux.
- apply (elements_nodup (s := Node)).
- apply elements_complete.
Qed.
(** Use the lemma to get specialized versions *)
Lemma list_lookup_nth_list {R : Semiring.type} (l : list (list R)) (k : Node) :
list_lookup [] elements l k = List.nth (index_of k) l [].
Proof. apply list_lookup_nth_gen. Qed.
(** Specialization of [list_lookup_nth_gen] for scalar lists. *)
Lemma list_lookup_nth_R {R : Semiring.type} (l : list R) (k : Node) :
list_lookup 0 elements l k = List.nth (index_of k) l 0.
Proof. apply list_lookup_nth_gen. Qed.
(** ** nth_map with custom defaults (in-bounds version) *)
Lemma nth_map_inbound : forall (A B : Type) (f : A -> B) (l : list A) (i : nat) (dA : A) (dB : B),
(i < List.length l)%nat ->
List.nth i (List.map f l) dB = f (List.nth i l dA).
Proof.
intros A B f l i dA dB Hbound.
(** Use nth_indep to switch default, then nth_map *)
assert (Heq_default : List.nth i (List.map f l) dB = List.nth i (List.map f l) (f dA)).
{ apply nth_default_indep. rewrite List.length_map. exact Hbound. }
rewrite Heq_default.
apply List.map_nth.
Qed.
(** ** fold_left add = fold_right add for the additive monoid *)
Lemma fold_left_add_acc {R : Semiring.type} : forall (l : list R) (a : R),
List.fold_left add l a = a + List.fold_right add 0 l.
Proof.
induction l as [|x l' IH]; intros a; simpl.
- rewrite addr0. reflexivity.
- rewrite IH.
(** (a + x) + fold_right add 0 l' = a + (x + fold_right add 0 l') *)
rewrite addA. reflexivity.
Qed.
(** [fold_left add] over a list equals [fold_right add 0]. *)
Lemma fold_left_add_fold_right_add {R : Semiring.type} : forall (l : list R),
List.fold_left add l 0 = List.fold_right add 0 l.
Proof.
intros l. rewrite fold_left_add_acc. rewrite add0r. reflexivity.
Qed.
(** ** combine distributes over map *)
Lemma combine_map : forall (A B C : Type) (f : A -> B) (g : A -> C) (l : list A),
List.combine (List.map f l) (List.map g l) = List.map (fun x => (f x, g x)) l.
Proof.
induction l as [|x l' IH]; simpl; [reflexivity |].
rewrite IH. reflexivity.
Qed.
(** ** dot_product of tabulated lists = sum of pointwise products *)
Lemma fold_right_add_map {R : Semiring.type} : forall (h : Node -> R),
List.fold_right add 0 (List.map h elements) =
List.fold_right (fun x y => h x + y) 0 elements.
Proof.
intro h.
induction elements as [|e es IH]; simpl; [reflexivity |].
f_equal. apply IH.
Qed.
(** Dot product of tabulated vectors equals the sum of pointwise products. *)
Lemma dot_product_map_eq_sum {R : Semiring.type} : forall (f g : Node -> R),
dot_product (List.map f elements) (List.map g elements) =
sum (fun x => f x * g x).
Proof.
intros f g.
unfold dot_product, sum.
rewrite combine_map.
rewrite List.map_map.
rewrite (fold_left_add_fold_right_add (List.map (fun x => f x * g x) elements)).
rewrite fold_right_add_map.
reflexivity.
Qed.
(** ** Re-tabulation lemma: map over elements reconstructs the list *)
Lemma list_lookup_tabulate : forall (A : Type) (d : A) (l : list A),
NoDup (elements (s := Node)) ->
List.length l = List.length (elements (s := Node)) ->
List.map (fun k => list_lookup d elements l k) (elements (s := Node)) = l.
Proof.
intros A d l Hnd Hlen.
revert l Hlen.
induction (elements (s := Node)) as [|e es IH]; intros l Hlen; simpl.
- destruct l; simpl in Hlen; try nia; reflexivity.
- destruct l as [|a l']; simpl in Hlen; [nia |].
simpl. f_equal.
+ unfold list_lookup, eq_decb.
destruct (fin_eq_dec e e) as [_ | Hneq]; [reflexivity | exfalso; apply Hneq; reflexivity].
+ inversion Hnd as [|? ? Hnin Hnd']; subst.
assert (Hext : forall z, In z es ->
(if eq_decb z e then a else list_lookup d es l' z) = list_lookup d es l' z).
{ intros z Hin_es. unfold eq_decb.
destruct (fin_eq_dec z e) as [Heq | Hneq]; [subst z; exfalso; apply Hnin; exact Hin_es | reflexivity]. }
rewrite (map_ext_in (fun z => if eq_decb z e then a else list_lookup d es l' z)
(fun z => list_lookup d es l' z) es Hext).
apply (IH Hnd' l').
nia.
Qed.
(** ** Helper lemmas for transpose_list via nth *)
Lemma nth_zip_with_cons_in_bounds : forall (A : Type) (xs : list A) (yss : list (list A))
(i : nat) (v : A),
(i < List.length xs)%nat -> (i < List.length yss)%nat ->
List.nth i (zip_with cons xs yss) [] =
List.cons (List.nth i xs v) (List.nth i yss []).
Proof.
intros A xs yss i v Hxs Hyss.
revert xs yss Hxs Hyss.
induction i as [|i IH]; intros xs yss Hxs Hyss.
- destruct xs as [|x xs]; [simpl in Hxs; nia |].
destruct yss as [|ys yss]; [simpl in Hyss; nia |].
simpl. reflexivity.
- destruct xs as [|x xs]; [simpl in Hxs; nia |].
destruct yss as [|ys yss]; [simpl in Hyss; nia |].
simpl. apply IH; simpl in Hxs, Hyss; nia.
Qed.
(** [nth] on [map] returns the default when the index is out of bounds. *)
Lemma nth_map_out_of_bounds : forall (A B : Type) (f : A -> B) (l : list A) (n : nat) (d : B),
List.length l <= n -> List.nth n (List.map f l) d = d.
Proof.
induction l as [|a l' IH]; intros n d Hle; simpl.
- destruct n; reflexivity.
- destruct n as [|n']; simpl.
+ simpl in Hle. exfalso. apply (Nat.nle_succ_0 _ Hle).
+ apply IH. apply le_S_n. exact Hle.
Qed.
(** [combine l []] is always []. *)
Lemma combine_nil_r : forall (A B : Type) (l : list A), List.combine l (@nil B) = [].
Proof.
induction l; simpl; auto.
Qed.
(** Dot product with an empty vector is zero. *)
Lemma dot_product_nil {R : Semiring.type} : forall (v : list R), dot_product v [] = 0.
Proof.
intros v. unfold dot_product. rewrite combine_nil_r. simpl. reflexivity.
Qed.
(** ** Key lemma: transpose_list swaps element access under nth *)
(** The unrestricted statement is false for ragged inputs such as [[]; [a]].
The intended transpose law is the rectangular, in-bounds form below. *)
(** Key lemma: transpose swaps indices under [nth] for rectangular matrices. *)
Lemma nth_transpose_swap {R : Semiring.type} : forall (L : list (list R)) (i j : nat),
L <> [] ->
(forall xs ys : list R, In xs L -> In ys L -> List.length xs = List.length ys) ->
(j < List.length L)%nat ->
(i < List.length (List.hd [] L))%nat ->
List.nth j (List.nth i (transpose_list L) ([] : list R)) 0 =
List.nth i (List.nth j L ([] : list R)) 0.
Proof.
intros L i j HLne Hrect Hj Hi.
generalize dependent i. generalize dependent j.
induction L as [|xsh L' IH]; [congruence | ].
destruct L' as [|xssth L'']; [|].
- (* singleton case: L = [xsh] *)
intros j Hj i Hi. cbn in Hj, Hi. cbn.
assert (Hj0 : j = 0%nat) by nia.
subst j; cbn. apply nth_0_map_singleton.
- (* multi-row case: L = xsh :: xssth :: L'' *)
intros j Hj i Hi. cbn in Hj, Hi. cbn [hd] in Hi. cbn [hd].
assert (Hne_tail : xssth :: L'' <> []) by congruence.
assert (Hrect_tail : forall xs ys : list R,
In xs (xssth :: L'') -> In ys (xssth :: L'') -> length xs = length ys).
{ intros xs ys Hx Hy. apply Hrect; [right; exact Hx | right; exact Hy]. }
assert (Hlen_eq : length xsh = length xssth).
{ apply Hrect with (xs := xsh) (ys := xssth);
[left; reflexivity | right; left; reflexivity]. }
assert (Hi_xssth : i < length xssth).
{ apply (Nat.lt_le_trans _ _ _ Hi). rewrite Hlen_eq. apply Nat.le_refl. }
specialize (IH Hne_tail Hrect_tail).
simpl (transpose_list (xsh :: xssth :: L'')).
(** shared helper lemmas *)
assert (Hpos_xsh : 0 < length xsh) by nia.
assert (Hpos_xssth : 0 < length (xssth :: L'')) by (cbn; nia).
assert (Hrect_full : forall xs ys : list R,
In xs (xsh :: xssth :: L'') -> In ys (xsh :: xssth :: L'') ->
length xs = length ys /\ 0 < length xs).
{ intros xs0 ys0 Hx Hy. split; [apply Hrect; assumption | ].
assert (Hlen_xs : length xs0 = length xsh).
{ apply Hrect with (ys := xsh); [exact Hx | left; reflexivity]. }
rewrite Hlen_xs; exact Hpos_xsh. }
assert (Hi_transpose : i < length (transpose_list (xssth :: L''))).
{ pose proof (transpose_length (A := R) (xssth :: L'') xsh
Hpos_xssth Hpos_xsh
(fun xs Hx ys Hy => Hrect_full xs ys Hx Hy)) as Htlen.
rewrite <- Htlen; exact Hi. }
erewrite nth_zip_with_cons_in_bounds with (v := 0);
[| exact Hi | exact Hi_transpose].
destruct j as [|j'].
+ (* j = 0 *)
cbn. reflexivity.
+ (* j = S j' *)
assert (Hj'_bound : j' < length (xssth :: L'')) by (cbn; lia).
cbn. destruct L'' as [|zssth L''']; cbn.
* (* inner single row *)
assert (Hj'0 : j' = 0%nat) by (cbn in Hj'_bound; lia).
subst j'. cbn. apply nth_0_map_singleton.
* (* inner multi-row *)
apply IH. exact Hj'_bound. exact Hi_xssth.
Qed.
(** Key lemma: dot_product v1 v2 = sum(λk. v1[k]*v2[k]) when
|v1| = |v2| = |elements| *)
Lemma dot_product_eq_sum {R : Semiring.type} : forall (v1 v2 : list R),
List.length v1 = List.length (elements (s := Node)) ->
List.length v2 = List.length (elements (s := Node)) ->
dot_product v1 v2 =
sum (fun k => list_lookup 0 elements v1 k * list_lookup 0 elements v2 k).
Proof.
intros v1 v2 Hlen1 Hlen2.
pose proof (elements_nodup (s := Node)) as Hnd.
rewrite <- (@list_lookup_tabulate R 0 v1 Hnd Hlen1) at 1.
rewrite <- (@list_lookup_tabulate R 0 v2 Hnd Hlen2) at 1.
apply dot_product_map_eq_sum.
Qed.
(** Key lemma: of_list (mul_list L1 L2) r c =
dot_product(row_r(L1), col_c(L2)) *)
Lemma of_list_mul_list_as_dot_product {R : Semiring.type} : forall (L1 L2 : list (list R)) (r c : Node),
of_list (mul_list L1 L2) r c =
dot_product (list_lookup [] elements L1 r)
(list_lookup [] elements (transpose_list L2) c).
Proof.
intros L1 L2 r c.
unfold of_list, mul_list.
rewrite list_lookup_nth_R.
rewrite list_lookup_nth_list.
rewrite list_lookup_nth_list.
rewrite list_lookup_nth_list.
set (i := index_of r).
set (j := index_of c).
set (f := fun (row : list R) =>
List.map (fun col : list R => dot_product row col) (transpose_list L2)).
destruct (Nat.lt_ge_cases i (List.length L1)) as [Hlt | Hge].
- (* i in bounds *)
rewrite (nth_map_inbound (list R) (list R) f L1 i [] [] Hlt).
unfold f.
set (row_r := List.nth i L1 []).
set (g := fun (col : list R) => dot_product row_r col).
destruct (Nat.lt_ge_cases j (List.length (transpose_list L2))) as [Hlt2 | Hge2].
+ (* j in bounds *)
rewrite (nth_map_inbound (list R) R g (transpose_list L2) j ([] : list R) 0 Hlt2).
unfold g, row_r. reflexivity.
+ (* j out of bounds *)
rewrite (List.nth_overflow (transpose_list L2) (n := j) ([] : list R)) by exact Hge2.
assert (HL : List.nth j (List.map g (transpose_list L2)) 0 = 0).
{ apply (nth_map_out_of_bounds (list R) R g (transpose_list L2) j 0 Hge2). }
assert (HR : dot_product row_r [] = 0).
{ apply dot_product_nil. }
rewrite HL. symmetry. exact HR.
- (* i out of bounds *)
rewrite (List.nth_overflow L1 (n := i) ([] : list R)) by exact Hge.
rewrite (nth_map_out_of_bounds (list R) (list R) f L1 i ([] : list R) Hge).
unfold dot_product. simpl. destruct j; reflexivity.
Qed.
(** Helper: for a list where all rows have equal length, transpose has that many rows *)
Lemma transpose_list_length_eq_n {R : Semiring.type} : forall (L : list (list R)) (m : nat),
L <> [] ->
(forall row, In row L -> length row = m) ->
length (transpose_list L) = m.
Proof.
intros L m Hne Hrow.
destruct L as [|xsh L']; [congruence |]; clear Hne.
revert xsh Hrow. induction L' as [|xssth L'' IH]; intros xsh Hrow.
- simpl (transpose_list [xsh]). rewrite List.length_map. apply Hrow. left; reflexivity.
- change (transpose_list (xsh :: xssth :: L''))
with (zip_with cons xsh (transpose_list (xssth :: L''))).
rewrite zip_with_length.
assert (Hxsh_len : length xsh = m).
{ apply Hrow. left; reflexivity. }
assert (Hxssth_len : length xssth = m).
{ apply Hrow. right; left; reflexivity. }
assert (IH_eq := IH xssth (fun row Hin => Hrow row (in_cons xsh row (xssth :: L'') Hin))).
replace (length (transpose_list (xssth :: L''))) with m by (symmetry; exact IH_eq).
rewrite Hxsh_len. apply Nat.min_id.
Qed.
(** Helper: each row of transpose_list L has length = length L (rectangular, nonempty) *)
Lemma transpose_row_len_eq {R : Semiring.type} : forall (L : list (list R)),
L <> [] ->
(forall xs ys : list R, In xs L -> In ys L -> length xs = length ys) ->
forall row, In row (transpose_list L) -> length row = length L.
Proof.
induction L as [|xsh L' IH]; [congruence |]; intros Hne Hrect row Hin.
destruct L' as [|xssth L''].
- (* L = [xsh]: transpose = map (fun y => [y]) xsh *)
cbn in Hin. apply in_map_iff in Hin. destruct Hin as (x & Hx & Hin_xsh).
subst row. cbn. reflexivity.
- (* L = xsh :: xssth :: L'': transpose = zip_with cons xsh T *)
change (transpose_list (xsh :: xssth :: L''))
with (zip_with cons xsh (transpose_list (xssth :: L''))) in Hin.
(** Any element of zip_with is at some index k *)
apply (In_nth (A := list R) _ _ ([] : list R)) in Hin.
destruct Hin as (k & Hk & Hrow).
assert (Hk_xsh : (k < length xsh)%nat).
{ rewrite zip_with_length in Hk. apply (Nat.lt_le_trans _ _ _ Hk (Nat.le_min_l _ _)). }
assert (Hk_T : (k < length (transpose_list (xssth :: L'')))%nat).
{ rewrite zip_with_length in Hk. apply (Nat.lt_le_trans _ _ _ Hk (Nat.le_min_r _ _)). }
subst row.
rewrite (nth_zip_with_cons_in_bounds R xsh (transpose_list (xssth :: L'')) k (0 : R) Hk_xsh Hk_T).
simpl. f_equal.
(** Need: length (nth k (transpose_list (xssth :: L'')) []) = length (xssth :: L'') *)
pose proof (nth_In (A := list R) (transpose_list (xssth :: L'')) ([] : list R) Hk_T) as Hin_T.
apply IH with (row := nth k (transpose_list (xssth :: L'')) []); auto.
+ intro Hc; congruence.
+ intros xs ys Hx Hy. apply Hrect; [right; exact Hx | right; exact Hy].
Qed.
(** ** list_lookup_transpose: transpose swaps key lookup *)
Lemma list_lookup_transpose {R : Semiring.type} : forall (L : list (list R)) (c y : Node),
(forall xs ys : list R, In xs L -> In ys L -> length xs = length ys) ->
list_lookup 0 elements (list_lookup [] elements (transpose_list L) c) y =
list_lookup 0 elements (list_lookup [] elements L y) c.
Proof.
intros L c y Hrect.
rewrite !list_lookup_nth_gen.
set (i := index_of c). set (j := index_of y).
destruct (Nat.eq_dec (length L) 0) as [Hlen0 | Hlen_pos].
{ (* L is empty *)
assert (HL : L = []) by (destruct L; cbn in Hlen0; [reflexivity | nia]).
subst L. cbn. destruct j; destruct i; reflexivity. }
(** L is nonempty *)
assert (HLne : L <> []) by (intro H; subst L; cbn in Hlen_pos; nia).
set (h := List.hd ([] : list R) L).
set (m := length h).
assert (Hrow_len : forall row, In row L -> length row = m).
{ intros row Hin. unfold m, h. apply Hrect with (ys := List.hd ([] : list R) L).
- exact Hin.
- destruct L; [congruence | left; reflexivity]. }
assert (HlenT : length (transpose_list L) = m).
{ apply (transpose_list_length_eq_n L m HLne). exact Hrow_len. }
destruct (Nat.lt_ge_cases i m) as [Hi_lt | Hi_ge].
{ (* i < m: row index in bounds for transpose_list L *)
destruct (Nat.lt_ge_cases j (length L)) as [Hj_lt | Hj_ge].
{ (* both in bounds: use nth_transpose_swap *)
apply nth_transpose_swap;
[exact HLne | exact Hrect | exact Hj_lt | unfold m, h; exact Hi_lt]. }
{ (* j out of bounds: LHS = 0, RHS = 0 *)
apply eq_trans with (y := 0).
2: { apply eq_sym. apply List.nth_overflow.
assert (Hnth_j_L : nth j L ([] : list R) = ([] : list R)).
{ apply List.nth_overflow. exact Hj_ge. }
rewrite Hnth_j_L. cbn. nia. }
apply List.nth_overflow.
assert (Hlen_row : length (nth i (transpose_list L) ([] : list R)) = length L).
{ eapply transpose_row_len_eq; eauto. eapply nth_In; rewrite HlenT; eauto. }
eapply Nat.le_trans; [| exact Hj_ge].
rewrite <- Hlen_row. apply Nat.le_refl. } }
{ (* i out of bounds: nth i (transpose_list L) = [] *)
assert (Hnth_i_T : nth i (transpose_list L) ([] : list R) = ([] : list R)).
{ apply List.nth_overflow. rewrite HlenT. exact Hi_ge. }
rewrite Hnth_i_T.
(** Prove nth j [] 0 = 0 by destructing j, then handle RHS *)
destruct j as [|j'].
{ simpl. apply eq_sym. apply List.nth_overflow.
assert (Hlen_row_j : length (nth 0 L ([] : list R)) = m).
{ apply Hrow_len. eapply nth_In; eauto; nia. }
rewrite Hlen_row_j. exact Hi_ge. }
{ simpl.
destruct (Nat.lt_ge_cases (S j') (length L)) as [Hj_lt' | Hj_ge'].
{ apply eq_sym. apply List.nth_overflow.
assert (Hlen_row_j : length (nth (S j') L ([] : list R)) = m).
{ apply Hrow_len. eapply nth_In; eauto. }
rewrite Hlen_row_j. exact Hi_ge. }
{ assert (Hnth_j_L : nth (S j') L ([] : list R) = ([] : list R)).
{ apply List.nth_overflow. exact Hj_ge'. }
rewrite Hnth_j_L. destruct i; reflexivity. } } }
Qed.
(** ** Main proof: of_list_mul_list_gen *)
Lemma of_list_mul_list_gen {R : Semiring.type} : forall (L1 L2 : list (list R)) (r c : Node),
List.length L1 = List.length (elements (s := Node)) ->
(forall row : list R, In row L1 ->
List.length row = List.length (elements (s := Node))) ->
List.length L2 = List.length (elements (s := Node)) ->
(forall row : list R, In row L2 ->
List.length row = List.length (elements (s := Node))) ->
of_list (mul_list L1 L2) r c = matrix_mul (of_list L1) (of_list L2) r c.
Proof.
intros L1 L2 r c HlenL1 HrowL1 HlenL2 HrowL2.
set (n := List.length (elements (s := Node))).
assert (HrectL1 : forall xs ys, In xs L1 -> In ys L1 -> length xs = length ys).
{ intros xs ys Hx Hy. rewrite (HrowL1 _ Hx), (HrowL1 _ Hy). reflexivity. }
assert (HrectL2 : forall xs ys, In xs L2 -> In ys L2 -> length xs = length ys).
{ intros xs ys Hx Hy. rewrite (HrowL2 _ Hx), (HrowL2 _ Hy). reflexivity. }
rewrite (of_list_mul_list_as_dot_product L1 L2 r c).
unfold matrix_mul, of_list.
rewrite (dot_product_eq_sum
(list_lookup [] elements L1 r)
(list_lookup [] elements (transpose_list L2) c)).
- apply sum_ext. intro y.
rewrite (list_lookup_transpose L2 c y); [reflexivity | exact HrectL2].
- (* Prove: length (list_lookup [] elements L1 r) = n *)
rewrite list_lookup_nth_gen.
set (i := index_of r).
assert (Hi : (i < List.length L1)%nat).
{ subst i. rewrite HlenL1. apply index_of_bound. }
apply nth_In with (d := [] : list R) in Hi as Hin.
apply HrowL1 in Hin. unfold n. exact Hin.
- (* Prove: length (list_lookup [] elements (transpose_list L2) c) = n *)
rewrite list_lookup_nth_gen.
set (i := index_of c).
(** Need: length (nth i (transpose_list L2) []) = n
Lemma: for n×n rectangular L2, each row of transpose has length = length L2 = n *)
assert (Htranspose_nth_len : length (nth i (transpose_list L2) []) = length L2).
{ (* i is in-bounds for transpose_list L2 since L2 is n×n *)
assert (L2_nonempty : L2 <> []).
{ intro H; subst L2. cbn in HlenL2.
pose proof (elements_complete (s := Node) r) as Hin_r.
destruct (elements (s := Node)); cbn in *; [inversion Hin_r | nia]. }
assert (Hi_bound : (i < length (transpose_list L2))%nat).
{ rewrite (transpose_list_length_eq_n L2 n L2_nonempty HrowL2). subst i. apply index_of_bound. }
pose proof (nth_In (A := list R) (transpose_list L2) ([] : list R) Hi_bound) as Hin_row.
apply transpose_row_len_eq with (row := nth i (transpose_list L2) []).
- exact L2_nonempty.
- exact HrectL2.
- exact Hin_row. }
rewrite Htranspose_nth_len. rewrite HlenL2. unfold n. reflexivity.
Qed.
(** Base case: [pow_list (to_list m) 0 = to_list I]. *)
Lemma pow_list_base {R : Semiring.type} : forall (m : @Matrix R),
pow_list (to_list m) 0 = to_list I.
Proof.
intros m.
unfold pow_list, to_list, I.
reflexivity.
Qed.
(** ** Helper lemmas: to_list and pow_list preserve the square shape *)
Lemma to_list_length {R : Semiring.type} : forall (m : @Matrix R),
length (to_list m) = length (elements (s := Node)).
Proof.
intros m. unfold to_list. rewrite List.length_map. reflexivity.
Qed.
(** Every row of [to_list m] has length [|elements|]. *)
Lemma to_list_row_length {R : Semiring.type} : forall (m : @Matrix R) (row : list R),
In row (to_list m) -> length row = length (elements (s := Node)).
Proof.
intros m row Hin.
unfold to_list in Hin.
rewrite in_map_iff in Hin. destruct Hin as (r & Hrow & Hin_r).
subst row. rewrite List.length_map. reflexivity.
Qed.
(** [mul_list] preserves the outer length (number of rows). *)
Lemma mul_list_length {R : Semiring.type} : forall (la lb : list (list R)),
length (mul_list la lb) = length la.
Proof.
intros la lb. unfold mul_list. rewrite List.length_map. reflexivity.
Qed.
(** For a square matrix, the transpose has the same dimension. *)
Lemma transpose_list_length_square {R : Semiring.type} : forall (lb : list (list R)),
length lb = length (elements (s := Node)) ->
(forall row, In row lb -> length row = length (elements (s := Node))) ->
length (transpose_list lb) = length (elements (s := Node)).
Proof.
intros lb Hlen Hrow.
destruct lb as [|lbh lbt].
- (* lb = [] *)
simpl. rewrite <- Hlen. reflexivity.
- (* lb is nonempty *)
apply (transpose_list_length_eq_n (lbh :: lbt) (length (elements (s := Node)))).