// Copyright (c) 2026 Petr Balvín (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) }