Files
tensor/linalg/mat_test.go
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

249 lines
7.4 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package linalg
import (
"sourcedock.dev/petrbalvin/tensor/internal/core"
"strings"
"testing"
)
func TestIdentity(t *testing.T) {
id, err := core.Identity(core.Int, 2)
if err != nil {
t.Fatalf("Identity: %v", err)
}
want := mustFromInts(t, []int64{1, 0, 0, 1}, 2, 2)
if !core.Equal(want, id) {
t.Fatalf("Identity int: %s", id)
}
fid, err := core.Identity(core.Float, 3)
if err != nil {
t.Fatalf("Identity float: %v", err)
}
if fid.Dtype() != core.Float {
t.Fatalf("Identity dtype: %s", fid.Dtype())
}
if v, _ := core.FloatAt(fid, 2, 2); v != 1 {
t.Fatalf("Identity corner: %v", v)
}
if _, err := core.Identity(core.Int, -1); err == nil || !strings.Contains(err.Error(), "zero or greater") {
t.Fatalf("Identity negative: %v", err)
}
}
func TestMatMul2D(t *testing.T) {
a := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6}, 2, 3)
b := mustFromInts(t, []int64{7, 8, 9, 10, 11, 12}, 3, 2)
c, err := core.MatMul2D(a, b)
if err != nil {
t.Fatalf("MatMul: %v", err)
}
if c.Dtype() != core.Int {
t.Fatalf("MatMul dtype: %s", c.Dtype())
}
want := mustFromInts(t, []int64{58, 64, 139, 154}, 2, 2)
if !core.Equal(want, c) {
t.Fatalf("MatMul: %s", c)
}
// Promotion and float values.
f := mustFromFloats(t, []float64{0.5, 1.5, 2.5}, 1, 3)
fc, err := core.MatMul2D(f, b)
if err != nil {
t.Fatalf("MatMul float: %v", err)
}
if fc.Dtype() != core.Float {
t.Fatalf("MatMul promote: %s", fc.Dtype())
}
// 0.5*7 + 1.5*9 + 2.5*11 = 44.5, 0.5*8 + 1.5*10 + 2.5*12 = 49
if v, _ := core.FloatAt(fc, 0, 0); v != 44.5 {
t.Fatalf("MatMul float value: %v", v)
}
if v, _ := core.FloatAt(fc, 0, 1); v != 49 {
t.Fatalf("MatMul float value: %v", v)
}
}
// TestMatMul2DFloat32MixedInt pins the mixed float32 x int product: an
// core.Int right operand used to reach the float32 kernel and slice b's nil
// floats32 payload, panicking instead of promoting.
func TestMatMul2DFloat32MixedInt(t *testing.T) {
a := mustFromFloat32s(t, []float32{1, 2, 3, 4}, 2, 2)
b := mustFromInts(t, []int64{5, 6, 7, 8}, 2, 2)
c, err := core.MatMul2D(a, b)
if err != nil {
t.Fatalf("MatMul float32 x int: %v", err)
}
if c.Dtype() != core.Float32 {
t.Fatalf("MatMul float32 x int dtype: %s", c.Dtype())
}
// [1*5+2*7, 1*6+2*8, 3*5+4*7, 3*6+4*8] = [19, 22, 43, 50]
want := mustFromFloat32s(t, []float32{19, 22, 43, 50}, 2, 2)
if !core.Equal(want, c) {
t.Fatalf("MatMul float32 x int: %s", c)
}
// The mirrored int x float32 product keeps working.
d, err := core.MatMul2D(b, a)
if err != nil {
t.Fatalf("MatMul int x float32: %v", err)
}
wantD := mustFromFloat32s(t, []float32{23, 34, 31, 46}, 2, 2)
if !core.Equal(wantD, d) {
t.Fatalf("MatMul int x float32: %s", d)
}
}
func TestMatMulVectorShapes(t *testing.T) {
m := mustFromInts(t, []int64{1, 2, 3, 4}, 2, 2)
v := mustFromInts(t, []int64{5, 6}, 2)
mv, err := core.MatMul2D(m, v)
if err != nil {
t.Fatalf("matrix × vector: %v", err)
}
// [1*5+2*6, 3*5+4*6] = [17, 39]
want := mustFromInts(t, []int64{17, 39}, 2)
if !core.Equal(want, mv) {
t.Fatalf("matrix × vector: %s", mv)
}
vm, err := core.MatMul2D(v, m)
if err != nil {
t.Fatalf("vector × matrix: %v", err)
}
// [5*1+6*3, 5*2+6*4] = [23, 34]
wantVM := mustFromInts(t, []int64{23, 34}, 2)
if !core.Equal(wantVM, vm) {
t.Fatalf("vector × matrix: %s", vm)
}
}
func TestMatMulErrors(t *testing.T) {
a := mustFromInts(t, []int64{1, 2, 3, 4}, 2, 2)
wrong2D := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6}, 3, 2)
if _, err := core.MatMul2D(a, wrong2D); err == nil || !strings.Contains(err.Error(), "inner dimensions must agree") {
t.Fatalf("inner mismatch: %v", err)
}
wrongV := mustFromInts(t, []int64{1, 2, 3}, 3)
if _, err := core.MatMul2D(a, wrongV); err == nil || !strings.Contains(err.Error(), "inner dimensions must agree") {
t.Fatalf("vector mismatch: %v", err)
}
if _, err := core.MatMul2D(wrongV, a); err == nil || !strings.Contains(err.Error(), "inner dimensions must agree") {
t.Fatalf("vector × matrix mismatch: %v", err)
}
v := mustFromInts(t, []int64{1}, 1)
if _, err := core.MatMul2D(v, v); err == nil || !strings.Contains(err.Error(), "unsupported shapes") {
t.Fatalf("1-D × 1-D is Dot: %v", err)
}
c3 := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6, 7, 8}, 2, 2, 2)
if _, err := core.MatMul2D(c3, c3); err == nil || !strings.Contains(err.Error(), "unsupported shapes") {
t.Fatalf("3-D matmul: %v", err)
}
}
func TestTranspose(t *testing.T) {
m := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6}, 2, 3)
tt := core.Transpose(m)
want := mustFromInts(t, []int64{1, 4, 2, 5, 3, 6}, 3, 2)
if !core.Equal(want, tt) {
t.Fatalf("Transpose: %s", tt)
}
// 1-D transposition is a copy.
v := mustFromInts(t, []int64{1, 2}, 2)
if !core.Equal(v, core.Transpose(v)) {
t.Fatalf("Transpose 1-D: %s", core.Transpose(v))
}
// 3-D reverses all dimensions.
c := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6, 7, 8}, 2, 2, 2)
ct := core.Transpose(c)
if shape := ct.Shape(); shape[0] != 2 || shape[1] != 2 || shape[2] != 2 {
t.Fatalf("Transpose 3-D shape: %v", shape)
}
if v, _ := core.IntAt(ct, 0, 0, 1); v != 5 {
t.Fatalf("Transpose 3-D value: %d", v)
}
}
func TestReshape(t *testing.T) {
a := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6}, 2, 3)
r, err := core.Reshape(a, 3, 2)
if err != nil {
t.Fatalf("Reshape: %v", err)
}
want := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6}, 3, 2)
if !core.Equal(want, r) {
t.Fatalf("Reshape: %s", r)
}
flat, err := core.Reshape(a, 6)
if err != nil || flat.Len() != 6 {
t.Fatalf("Reshape flat: %s %v", flat, err)
}
if _, err := core.Reshape(a, 4); err == nil || !strings.Contains(err.Error(), "do not fill the shape") {
t.Fatalf("Reshape count: %v", err)
}
if _, err := core.Reshape(a); err == nil || !strings.Contains(err.Error(), "at least one dimension") {
t.Fatalf("Reshape empty: %v", err)
}
}
func TestRowCol(t *testing.T) {
m := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6}, 2, 3)
row, err := core.Row(m, 1)
if err != nil {
t.Fatalf("Row: %v", err)
}
wantRow := mustFromInts(t, []int64{4, 5, 6}, 3)
if !core.Equal(wantRow, row) {
t.Fatalf("Row: %s", row)
}
col, err := core.Col(m, 1)
if err != nil {
t.Fatalf("Col: %v", err)
}
wantCol := mustFromInts(t, []int64{2, 5}, 2)
if !core.Equal(wantCol, col) {
t.Fatalf("Col: %s", col)
}
if _, err := core.Row(m, 2); err == nil || !strings.Contains(err.Error(), "out of range") {
t.Fatalf("Row range: %v", err)
}
if _, err := core.Col(m, 3); err == nil || !strings.Contains(err.Error(), "out of range") {
t.Fatalf("Col range: %v", err)
}
v := mustFromInts(t, []int64{1}, 1)
if _, err := core.Row(v, 0); err == nil || !strings.Contains(err.Error(), "needs a 2-D array") {
t.Fatalf("Row 1-D: %v", err)
}
if _, err := core.Col(v, 0); err == nil || !strings.Contains(err.Error(), "needs a 2-D array") {
t.Fatalf("Col 1-D: %v", err)
}
}
// TestKronIntExact pins the int Kronecker product: products above 2^53
// used to round-trip through float64 and lose their low bits.
func TestKronIntExact(t *testing.T) {
big := int64(1) << 53
a := mustFromInts(t, []int64{big + 1}, 1, 1)
b := mustFromInts(t, []int64{2}, 1, 1)
got, err := core.Kron(a, b)
if err != nil {
t.Fatal(err)
}
if got.Dtype() != core.Int {
t.Fatalf("Kron int dtype: %s", got.Dtype())
}
if v, _ := core.IntAt(got, 0, 0); v != (big+1)*2 {
t.Fatalf("Kron int exact: got %d, want %d", v, (big+1)*2)
}
}