// Copyright (c) 2026 Petr Balvín (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) } }