249 lines
7.4 KiB
Go
249 lines
7.4 KiB
Go
// 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)
|
||
}
|
||
}
|