172 lines
5.1 KiB
Go
172 lines
5.1 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (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
|
||
|
|
}
|