symbolic mathematics engine in OCaml with differentiation, integration, simplification, and numerical methods
17

Configure Feed

Select the types of activity you want to include in your feed.

leibniz / lib / matrix.ml
3.9 kB 137 lines
1open Expr 2open Simplify 3 4type matrix = expr array array 5 6let create rows cols init = 7 Array.init rows (fun i -> Array.init cols (fun j -> init i j)) 8 9let identity n = 10 create n n (fun i j -> if i = j then Const 1.0 else Const 0.0) 11 12let rows m = Array.length m 13let cols m = if Array.length m > 0 then Array.length m.(0) else 0 14 15let get m i j = m.(i).(j) 16let set m i j v = m.(i).(j) <- v 17 18let map f m = 19 Array.map (Array.map f) m 20 21let transpose m = 22 let r = rows m in 23 let c = cols m in 24 create c r (fun i j -> m.(j).(i)) 25 26let add m1 m2 = 27 if rows m1 <> rows m2 || cols m1 <> cols m2 then 28 failwith "matrix dimensions must match for addition" 29 else 30 create (rows m1) (cols m1) (fun i j -> 31 simplify (Add (m1.(i).(j), m2.(i).(j))) 32 ) 33 34let mult m1 m2 = 35 if cols m1 <> rows m2 then 36 failwith "incompatible matrix dimensions for multiplication" 37 else 38 create (rows m1) (cols m2) (fun i j -> 39 let sum = ref (Const 0.0) in 40 for k = 0 to cols m1 - 1 do 41 sum := Add (!sum, Mul (m1.(i).(k), m2.(k).(j))) 42 done; 43 simplify !sum 44 ) 45 46let scalar_mult s m = 47 map (fun e -> simplify (Mul (s, e))) m 48 49let det m = 50 let n = rows m in 51 if n <> cols m then failwith "determinant requires square matrix"; 52 53 let rec determinant mat size = 54 if size = 1 then mat.(0).(0) 55 else if size = 2 then 56 simplify (Sub (Mul (mat.(0).(0), mat.(1).(1)), 57 Mul (mat.(0).(1), mat.(1).(0)))) 58 else 59 let result = ref (Const 0.0) in 60 for j = 0 to size - 1 do 61 let minor = create (size - 1) (size - 1) (fun i k -> 62 let mi = if i < 0 then i else i + 1 in 63 let mk = if k < j then k else k + 1 in 64 mat.(mi).(mk) 65 ) in 66 let cofactor = determinant minor (size - 1) in 67 let sign = if j mod 2 = 0 then Const 1.0 else Const (-1.0) in 68 result := Add (!result, Mul (Mul (sign, mat.(0).(j)), cofactor)) 69 done; 70 simplify !result 71 in 72 determinant m n 73 74let inverse m = 75 let n = rows m in 76 if n <> cols m then None 77 else 78 let d = det m in 79 match d with 80 | Const 0.0 -> None 81 | _ -> 82 let adj = create n n (fun i j -> 83 let minor = create (n - 1) (n - 1) (fun mi mj -> 84 let si = if mi < i then mi else mi + 1 in 85 let sj = if mj < j then mj else mj + 1 in 86 m.(si).(sj) 87 ) in 88 let minor_det = det minor in 89 let sign = if (i + j) mod 2 = 0 then Const 1.0 else Const (-1.0) in 90 simplify (Mul (sign, minor_det)) 91 ) in 92 let adj_t = transpose adj in 93 Some (map (fun e -> simplify (Div (e, d))) adj_t) 94 95let trace m = 96 let n = min (rows m) (cols m) in 97 let sum = ref (Const 0.0) in 98 for i = 0 to n - 1 do 99 sum := Add (!sum, m.(i).(i)) 100 done; 101 simplify !sum 102 103let eigenvalues _m = 104 [] 105 106let rank m = 107 let r = rows m in 108 let c = cols m in 109 let temp = Array.map Array.copy m in 110 let rec count_pivots row col rank = 111 if row >= r || col >= c then rank 112 else 113 match temp.(row).(col) with 114 | Const 0.0 -> 115 let rec find_pivot i = 116 if i >= r then None 117 else match temp.(i).(col) with 118 | Const 0.0 -> find_pivot (i + 1) 119 | _ -> Some i 120 in 121 (match find_pivot (row + 1) with 122 | None -> count_pivots row (col + 1) rank 123 | Some i -> 124 let tmp_row = temp.(row) in 125 temp.(row) <- temp.(i); 126 temp.(i) <- tmp_row; 127 count_pivots row col rank) 128 | pivot -> 129 for i = row + 1 to r - 1 do 130 let factor = Div (temp.(i).(col), pivot) in 131 for j = col to c - 1 do 132 temp.(i).(j) <- simplify (Sub (temp.(i).(j), Mul (factor, temp.(row).(j)))) 133 done 134 done; 135 count_pivots (row + 1) (col + 1) (rank + 1) 136 in 137 count_pivots 0 0 0