344 lines
10 KiB
Go
344 lines
10 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
|
|
// SPDX-License-Identifier: MIT
|
|||
|
|
|
|||
|
|
package linalg
|
|||
|
|
|
|||
|
|
import (
|
|||
|
|
"math"
|
|||
|
|
|
|||
|
|
"sourcedock.dev/petrbalvin/tensor/internal/base"
|
|||
|
|
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
|||
|
|
)
|
|||
|
|
|
|||
|
|
// Linear algebra. Every entry point requires a square
|
|||
|
|
// 2-D matrix. One LU decomposition with partial pivoting backs Det,
|
|||
|
|
// Solve and Inv; the same kernel serves both real and complex element
|
|||
|
|
// types through generics, so complex systems solve exactly where their
|
|||
|
|
// real counterparts do. A singular matrix is an error for Solve and Inv,
|
|||
|
|
// while the determinants return 0. Results are float64/complex128
|
|||
|
|
// approximations, as is standard for elimination with partial pivoting;
|
|||
|
|
// near-singular matrices lose precision without failing.
|
|||
|
|
//
|
|||
|
|
// Type summary per the ladder: Solve and Inv promote to complex128 when
|
|||
|
|
// any operand is complex; Det and Trace answer only real matrices and
|
|||
|
|
// point complex callers at DetComplex and TraceComplex.
|
|||
|
|
|
|||
|
|
// squareFloatMatrix converts a real array to a square float64 matrix,
|
|||
|
|
// promoting ints and float32 on the way; it errors with the caller's
|
|||
|
|
// name when the shape is wrong. The rows are views over one flat
|
|||
|
|
// backing slice, so the conversion costs two allocations. The rows are
|
|||
|
|
// handed to the solver, which may reorder the row headers in place, so
|
|||
|
|
// callers must treat the result as consumed.
|
|||
|
|
func squareFloatMatrix(a *core.Array, name string) ([][]float64, error) {
|
|||
|
|
if a.NDim() != 2 || a.Shape()[0] != a.Shape()[1] {
|
|||
|
|
return nil, base.Errf("%s: needs a square 2-D matrix, got shape %s", name, base.ShapeText(a.Shape()))
|
|||
|
|
}
|
|||
|
|
n := a.Shape()[0]
|
|||
|
|
back := make([]float64, n*n)
|
|||
|
|
if a.Dtype() == core.Float && !a.Strided() && len(a.RawFloats()) == n*n {
|
|||
|
|
copy(back, a.RawFloats())
|
|||
|
|
} else {
|
|||
|
|
for i := range n * n {
|
|||
|
|
back[i] = a.FloatAt(i)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
rows := make([][]float64, n)
|
|||
|
|
for i := range n {
|
|||
|
|
rows[i] = back[i*n : (i+1)*n]
|
|||
|
|
}
|
|||
|
|
return rows, nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// squareComplexMatrix converts a complex array to a square matrix of
|
|||
|
|
// complex128 values with the same shape contract as its real twin.
|
|||
|
|
func squareComplexMatrix(a *core.Array, name string) ([][]complex128, error) {
|
|||
|
|
if a.Dtype() != core.Complex {
|
|||
|
|
return nil, base.Errf("%s: needs a complex matrix", name)
|
|||
|
|
}
|
|||
|
|
if a.NDim() != 2 || a.Shape()[0] != a.Shape()[1] {
|
|||
|
|
return nil, base.Errf("%s: needs a square 2-D matrix, got shape %s", name, base.ShapeText(a.Shape()))
|
|||
|
|
}
|
|||
|
|
n := a.Shape()[0]
|
|||
|
|
back := make([]complex128, n*n)
|
|||
|
|
if !a.Strided() && len(a.RawComplexes()) == n*n {
|
|||
|
|
copy(back, a.RawComplexes())
|
|||
|
|
} else {
|
|||
|
|
for i := range n * n {
|
|||
|
|
back[i] = a.ComplexAt(i)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
rows := make([][]complex128, n)
|
|||
|
|
for i := range n {
|
|||
|
|
rows[i] = back[i*n : (i+1)*n]
|
|||
|
|
}
|
|||
|
|
return rows, nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// identityColumns builds the columns of the n×n identity, the standard
|
|||
|
|
// right-hand side of the inverse. The columns slice into one flat
|
|||
|
|
// backing array: identical values, two allocations instead of n+1.
|
|||
|
|
func identityColumns[T scalar](n int) [][]T {
|
|||
|
|
back := make([]T, n*n)
|
|||
|
|
cols := make([][]T, n)
|
|||
|
|
var one T = 1
|
|||
|
|
for j := range n {
|
|||
|
|
col := back[j*n : (j+1)*n]
|
|||
|
|
col[j] = one
|
|||
|
|
cols[j] = col
|
|||
|
|
}
|
|||
|
|
return cols
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// requireReal refuses the complex dtype at an entry point that reads
|
|||
|
|
// its inputs as real numbers. Without the gate the first FloatAt would
|
|||
|
|
// fall through to the complex array's empty int payload and index out
|
|||
|
|
// of range, which is a panic rather than the error the caller expects.
|
|||
|
|
func requireReal(name string, arrays ...*core.Array) error {
|
|||
|
|
for _, a := range arrays {
|
|||
|
|
if a.Dtype() == core.Complex {
|
|||
|
|
return base.Errf("%s: complex arrays are not supported", name)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
return nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// Det returns the determinant of a square real matrix (ints and
|
|||
|
|
// float32 promote). A singular matrix yields 0, not an error; complex
|
|||
|
|
// matrices are answered by DetComplex.
|
|||
|
|
func Det(a *core.Array) (float64, error) {
|
|||
|
|
if a.Dtype() == core.Complex {
|
|||
|
|
return 0, base.Errf("Det: complex matrices are answered by DetComplex")
|
|||
|
|
}
|
|||
|
|
m, err := squareFloatMatrix(a, "Det")
|
|||
|
|
if err != nil {
|
|||
|
|
return 0, err
|
|||
|
|
}
|
|||
|
|
_, parity := base.Factor(m)
|
|||
|
|
det := 1.0 * float64(parity)
|
|||
|
|
for i := range m {
|
|||
|
|
det *= m[i][i]
|
|||
|
|
}
|
|||
|
|
return det, nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// DetComplex returns the determinant of a square complex matrix as a
|
|||
|
|
// complex128 value. A singular matrix yields 0, not an error.
|
|||
|
|
func DetComplex(a *core.Array) (complex128, error) {
|
|||
|
|
m, err := squareComplexMatrix(a, "DetComplex")
|
|||
|
|
if err != nil {
|
|||
|
|
return 0, err
|
|||
|
|
}
|
|||
|
|
_, parity := base.Factor(m)
|
|||
|
|
det := complex(float64(parity), 0)
|
|||
|
|
for i := range m {
|
|||
|
|
det *= m[i][i]
|
|||
|
|
}
|
|||
|
|
return det, nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// solveRHS converts b into float columns of length n. b may be a vector
|
|||
|
|
// of length n or a matrix with n rows.
|
|||
|
|
func solveRHS(b *core.Array, n int) ([][]float64, error) {
|
|||
|
|
switch {
|
|||
|
|
case b.NDim() == 1 && b.Len() == n:
|
|||
|
|
col := make([]float64, n)
|
|||
|
|
if b.Dtype() == core.Float && !b.Strided() && len(b.RawFloats()) == n {
|
|||
|
|
copy(col, b.RawFloats())
|
|||
|
|
} else {
|
|||
|
|
for i := range n {
|
|||
|
|
col[i] = b.FloatAt(i)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
return [][]float64{col}, nil
|
|||
|
|
case b.NDim() == 2 && b.Shape()[0] == n:
|
|||
|
|
w := b.Shape()[1]
|
|||
|
|
cols := make([][]float64, w)
|
|||
|
|
dense := b.Dtype() == core.Float && !b.Strided() && len(b.RawFloats()) == n*w
|
|||
|
|
for j := range w {
|
|||
|
|
col := make([]float64, n)
|
|||
|
|
for i := range n {
|
|||
|
|
if dense {
|
|||
|
|
col[i] = b.RawFloats()[i*w+j]
|
|||
|
|
} else {
|
|||
|
|
col[i] = b.FloatAt(i*w + j)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
cols[j] = col
|
|||
|
|
}
|
|||
|
|
return cols, nil
|
|||
|
|
}
|
|||
|
|
return nil, base.Errf("Solve: b must be a vector of length %d or a matrix with %d rows, got shape %s",
|
|||
|
|
n, n, base.ShapeText(b.Shape()))
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// complexRHS converts b into complex columns of length n under the same
|
|||
|
|
// shape contract as solveRHS; a real right-hand side promotes wholesale.
|
|||
|
|
func complexRHS(b *core.Array, n int) ([][]complex128, error) {
|
|||
|
|
switch {
|
|||
|
|
case b.NDim() == 1 && b.Len() == n:
|
|||
|
|
col := make([]complex128, n)
|
|||
|
|
if b.Dtype() == core.Complex && !b.Strided() && len(b.RawComplexes()) == n {
|
|||
|
|
copy(col, b.RawComplexes())
|
|||
|
|
} else {
|
|||
|
|
for i := range n {
|
|||
|
|
col[i] = b.ComplexAt(i)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
return [][]complex128{col}, nil
|
|||
|
|
case b.NDim() == 2 && b.Shape()[0] == n:
|
|||
|
|
w := b.Shape()[1]
|
|||
|
|
cols := make([][]complex128, w)
|
|||
|
|
dense := b.Dtype() == core.Complex && !b.Strided() && len(b.RawComplexes()) == n*w
|
|||
|
|
for j := range w {
|
|||
|
|
col := make([]complex128, n)
|
|||
|
|
for i := range n {
|
|||
|
|
if dense {
|
|||
|
|
col[i] = b.RawComplexes()[i*w+j]
|
|||
|
|
} else {
|
|||
|
|
col[i] = b.ComplexAt(i*w + j)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
cols[j] = col
|
|||
|
|
}
|
|||
|
|
return cols, nil
|
|||
|
|
}
|
|||
|
|
return nil, base.Errf("Solve: b must be a vector of length %d or a matrix with %d rows, got shape %s",
|
|||
|
|
n, n, base.ShapeText(b.Shape()))
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// columnsToFloatArray assembles solved columns back into a row-major
|
|||
|
|
// array shaped like the original right-hand side: like[0] rows in both
|
|||
|
|
// the vector and the matrix case.
|
|||
|
|
func columnsToFloatArray(cols [][]float64, like []int) *core.Array {
|
|||
|
|
rows := like[0]
|
|||
|
|
out := core.New(core.Float, like...)
|
|||
|
|
for i := range rows {
|
|||
|
|
base := i * len(cols)
|
|||
|
|
for c, col := range cols {
|
|||
|
|
out.RawFloats()[base+c] = col[i]
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
return out
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// columnsToComplexArray is the complex twin of columnsToFloatArray.
|
|||
|
|
func columnsToComplexArray(cols [][]complex128, like []int) *core.Array {
|
|||
|
|
out := core.New(core.Complex, like...)
|
|||
|
|
rows := like[0]
|
|||
|
|
for r := range rows {
|
|||
|
|
for c, col := range cols {
|
|||
|
|
out.RawComplexes()[r*len(cols)+c] = col[r]
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
return out
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// Solve returns x such that a·x = b, where b is a vector of length n or
|
|||
|
|
// a matrix with n rows. Real inputs give a float result; any complex
|
|||
|
|
// operand promotes the whole system to complex128 per the ladder.
|
|||
|
|
func Solve(a, b *core.Array) (*core.Array, error) {
|
|||
|
|
if a.Dtype() == core.Complex || b.Dtype() == core.Complex {
|
|||
|
|
// Any complex operand promotes the whole system: a real A
|
|||
|
|
// widens alongside a complex right-hand side.
|
|||
|
|
if a.Dtype() != core.Complex {
|
|||
|
|
conv, cerr := core.Astype(a, core.Complex)
|
|||
|
|
if cerr != nil {
|
|||
|
|
return nil, cerr
|
|||
|
|
}
|
|||
|
|
a = conv
|
|||
|
|
}
|
|||
|
|
ac, err := squareComplexMatrix(a, "Solve")
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
rhs, err := complexRHS(b, len(ac))
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
x, err := base.SolveSystem("Solve", ac, rhs)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
return columnsToComplexArray(x, b.Shape()), nil
|
|||
|
|
}
|
|||
|
|
m, err := squareFloatMatrix(a, "Solve")
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
rhs, err := solveRHS(b, len(m))
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
x, err := base.SolveSystem("Solve", m, rhs)
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
return columnsToFloatArray(x, b.Shape()), nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// Inv returns the inverse of a square matrix. Real inputs give a float
|
|||
|
|
// array; a complex matrix answers in complex128.
|
|||
|
|
func Inv(a *core.Array) (*core.Array, error) {
|
|||
|
|
if a.Dtype() == core.Complex {
|
|||
|
|
ac, err := squareComplexMatrix(a, "Inv")
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
x, err := base.SolveSystem("Inv", ac, identityColumns[complex128](len(ac)))
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
return columnsToComplexArray(x, []int{len(ac), len(ac)}), nil
|
|||
|
|
}
|
|||
|
|
m, err := squareFloatMatrix(a, "Inv")
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
x, err := base.SolveSystem("Inv", m, identityColumns[float64](len(m)))
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
return columnsToFloatArray(x, []int{len(m), len(m)}), nil
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
// FitPolynomial fits the coefficients (lowest power first) of a
|
|||
|
|
// degree-th polynomial through (x, y) samples by QR least squares over
|
|||
|
|
// the Vandermonde matrix, which stays stable where the normal
|
|||
|
|
// equations would square the condition number.
|
|||
|
|
func FitPolynomial(x, y *core.Array, degree int) (*core.Array, error) {
|
|||
|
|
n := x.Len()
|
|||
|
|
if err := requireReal("FitPolynomial", x, y); err != nil {
|
|||
|
|
return nil, err
|
|||
|
|
}
|
|||
|
|
if degree < 0 {
|
|||
|
|
return nil, base.Errf("FitPolynomial: degree %d is negative", degree)
|
|||
|
|
}
|
|||
|
|
// degree+1 would wrap at the int maximum and slip past the sample
|
|||
|
|
// count check below as a negative.
|
|||
|
|
if degree == math.MaxInt {
|
|||
|
|
return nil, base.Errf("FitPolynomial: degree %d is too large", degree)
|
|||
|
|
}
|
|||
|
|
if n < degree+1 || y.Len() != n {
|
|||
|
|
return nil, base.Errf("FitPolynomial: need n ≥ degree+1 matching samples")
|
|||
|
|
}
|
|||
|
|
aMat, err := zeros(core.Float, []int{n, degree + 1})
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, base.Errf("FitPolynomial: %w", err)
|
|||
|
|
}
|
|||
|
|
for r := range n {
|
|||
|
|
pow := 1.0
|
|||
|
|
for d := range degree + 1 {
|
|||
|
|
aMat.SetFloatAt(r*(degree+1)+d, pow)
|
|||
|
|
pow *= x.FloatAt(r)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
yFlat, err := zeros(core.Float, []int{n})
|
|||
|
|
if err != nil {
|
|||
|
|
return nil, base.Errf("FitPolynomial: %w", err)
|
|||
|
|
}
|
|||
|
|
for i := range n {
|
|||
|
|
yFlat.SetFloatAt(i, y.FloatAt(i))
|
|||
|
|
}
|
|||
|
|
return LeastSquares(aMat, yFlat)
|
|||
|
|
}
|