mangrovedigital/tide-engine-api
0
1"""Minimal pure-Python linear algebra: ordinary/ridge least squares.
2
3solve_lstsq(X, y) returns the coefficients b minimising ||X b - y||^2 via the
4normal equations (X'X + lam I) b = X'y, solved with Gaussian elimination and
5partial pivoting. A tiny ridge term keeps near-collinear designs stable.
6"""
7
8
9def _matT_mat(X):
10 n = len(X)
11 m = len(X[0])
12 A = [[0.0] * m for _ in range(m)]
13 for r in range(n):
14 row = X[r]
15 for i in range(m):
16 xi = row[i]
17 if xi == 0.0:
18 continue
19 Ai = A[i]
20 for j in range(i, m):
21 Ai[j] += xi * row[j]
22 for i in range(m):
23 for j in range(i):
24 A[i][j] = A[j][i]
25 return A
26
27
28def _matT_vec(X, y):
29 m = len(X[0])
30 b = [0.0] * m
31 for r in range(len(X)):
32 row = X[r]
33 yr = y[r]
34 for i in range(m):
35 b[i] += row[i] * yr
36 return b
37
38
39def _solve(A, b):
40 m = len(A)
41 M = [row[:] + [b[i]] for i, row in enumerate(A)]
42 for col in range(m):
43 piv = max(range(col, m), key=lambda r: abs(M[r][col]))
44 if abs(M[piv][col]) < 1e-15:
45 raise ValueError("singular system")
46 M[col], M[piv] = M[piv], M[col]
47 pv = M[col][col]
48 for r in range(m):
49 if r == col:
50 continue
51 f = M[r][col] / pv
52 if f == 0.0:
53 continue
54 Mr, Mc = M[r], M[col]
55 for c in range(col, m + 1):
56 Mr[c] -= f * Mc[c]
57 return [M[i][m] / M[i][i] for i in range(m)]
58
59
60def solve_lstsq(X, y, ridge=1e-8):
61 """X: list of rows (each list of floats). y: list of floats. -> coeffs."""
62 A = _matT_mat(X)
63 b = _matT_vec(X, y)
64 m = len(A)
65 if ridge:
66 # scale ridge by mean diagonal so it is relative, not absolute
67 d = sum(A[i][i] for i in range(m)) / m
68 lam = ridge * (d if d > 0 else 1.0)
69 for i in range(m):
70 A[i][i] += lam
71 return _solve(A, b)
72
73
74def predict(X, coeffs):
75 return [sum(row[i] * coeffs[i] for i in range(len(coeffs))) for row in X]
76
77
78def rmse(a, b):
79 n = len(a)
80 return (sum((a[i] - b[i]) ** 2 for i in range(n)) / n) ** 0.5
81
82
83def mae(a, b):
84 n = len(a)
85 return sum(abs(a[i] - b[i]) for i in range(n)) / n
86 