symbolic mathematics engine in OCaml with differentiation, integration, simplification, and numerical methods
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