// Copyright (c) 2026 Petr BalvĂ­n (https://petrbalvin.org) // SPDX-License-Identifier: MIT package core import "sourcedock.dev/petrbalvin/tensor/internal/base" // Matrix utilities: Trace, Diagonal, Kron. These are the // smaller linear-algebra helpers that sit alongside MatMul / Solve / // Inv. Trace reduces a 2-D matrix along the main diagonal; Diagonal // extracts one diagonal as a 1-D array (offset > 0 picks a super- // diagonal, offset < 0 picks a sub-diagonal); Kron computes the // Kronecker product of two 2-D matrices, with a result whose shape is // the outer product of the inputs' shapes. // Trace returns the sum along the main diagonal of a 2-D square // matrix. The result is always float64: int arrays promote, complex // matrices are answered by TraceComplex. The narrow integer widths and // bool accumulate in int64 from exact widenings before the single // widening into the float64 result, the exactness rule the scalar // reductions carry. func Trace(a *Array) (float64, error) { if a.NDim() != 2 || a.Shape()[0] != a.Shape()[1] { return 0, errf("Trace: needs a square 2-D matrix, got shape %s", shapeText(a.Shape())) } if a.Dtype() == Complex { return 0, errf("Trace: complex matrices are answered by TraceComplex") } n := a.Shape()[0] if narrowRefused(a.Dtype()) { var is int64 for i := range n { is += a.intAt(i*n + i) } return float64(is), nil } var s float64 for i := range n { s += a.FloatAt(i*n + i) } return s, nil } // TraceComplex returns the main-diagonal sum of a complex square matrix. func TraceComplex(a *Array) (complex128, error) { if a.NDim() != 2 || a.Shape()[0] != a.Shape()[1] { return 0, errf("TraceComplex: needs a square 2-D matrix, got shape %s", shapeText(a.Shape())) } if a.Dtype() != Complex { return 0, errf("TraceComplex: needs a complex matrix") } var s complex128 n := a.Shape()[0] for i := range n { s += a.ComplexAt(i*n + i) } return s, nil } // Diagonal returns the elements on the offset-th diagonal of a 2-D // matrix as a 1-D array. offset = 0 is the main diagonal, offset > 0 a // super-diagonal, offset < 0 a sub-diagonal. The result's dtype follows // the input. func Diagonal(a *Array, offset int) (*Array, error) { if a.NDim() != 2 { return nil, errf("Diagonal: needs a 2-D array, got shape %s", shapeText(a.Shape())) } rows, cols := a.Shape()[0], a.Shape()[1] var startRow, startCol int var length int switch { case offset >= 0: startRow = 0 startCol = offset length = min(rows, cols-offset) default: startRow = -offset startCol = 0 length = min(rows+offset, cols) } if length <= 0 { return nil, errf("Diagonal: offset %d is out of range for shape %s", offset, shapeText(a.Shape())) } out := &Array{shape: []int{length}, dt: a.Dtype()} out.alloc(length) for i := range length { out.setFrom(i, a, (startRow+i)*cols+(startCol+i)) } return out, nil } // Kron returns the Kronecker product of two 2-D matrices. For shapes // (m, n) and (p, q) the result has shape (m*p, n*q) and element // (i*p+k, j*q+l) = a[i,j] * b[k,l]. The result dtype follows the // promotion ladder: int with int stays int, any float or complex // operand promotes. func Kron(a, b *Array) (*Array, error) { if a.NDim() != 2 || b.NDim() != 2 { return nil, errf("Kron: both inputs must be 2-D, got shapes %s and %s", shapeText(a.shape), shapeText(b.shape)) } for _, op := range []*Array{a, b} { if narrowRefused(op.Dtype()) { return nil, errf("Kron: dtype %s is not supported; convert with Astype", op.Dtype()) } } ma, na := a.Shape()[0], a.Shape()[1] mb, nb := b.Shape()[0], b.Shape()[1] dt := promote(a.Dtype(), b.Dtype()) // The outer-product shape must go through the same validation as // every other shape: unchecked products wrap on oversized inputs. total, outShape, terr := checkedDims([]int{ma * mb, na * nb}) if terr != nil { return nil, base.WrapErr("Kron", terr) } out := &Array{shape: outShape, dt: dt} out.alloc(total) if dt == Complex { for i := range ma { for j := range na { for k := range mb { for l := range nb { out.complexes[((i*mb+k)*(na*nb))+j*nb+l] = a.ComplexAt(i*na+j) * b.ComplexAt(k*nb+l) } } } } return out, nil } if dt == Int { // int x int stays exact: products must not round-trip through // float64, which loses low bits above 2^53. for i := range ma { for j := range na { for k := range mb { for l := range nb { outRow := (i*mb + k) * (na * nb) outCol := j*nb + l out.ints[outRow+outCol] = a.ints[i*na+j] * b.ints[k*nb+l] } } } } return out, nil } for i := range ma { for j := range na { av := a.FloatAt(i*na + j) for k := range mb { for l := range nb { // Both operands widen to float64, the product // narrows once per element, so an int // operand loses at most one rounding, like every // sibling kernel. prod := av * b.FloatAt(k*nb+l) pos := (i*mb+k)*(na*nb) + j*nb + l switch out.dt { case Float16: out.halves[pos] = HalfFromFloat64(prod) case Float32: out.floats32[pos] = float32(prod) default: out.floats[pos] = prod } } } } } return out, nil }