Files
tensor/internal/core/matrix2.go
T
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

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
}