81 lines
3.0 KiB
Go
81 lines
3.0 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|
// SPDX-License-Identifier: MIT
|
|
|
|
package core
|
|
|
|
import (
|
|
"math"
|
|
"testing"
|
|
)
|
|
|
|
// TestEinsumReductionKeepsFloat32WithEngine pins the routing of the
|
|
// single-operand axis reduction for a float32 operand: the fold declines
|
|
// it, because the fold widens, sums in float64 and narrows once, while
|
|
// the engine narrows into the float32 slot once per addend. The operand
|
|
// carries 1, 1, 1e8 and -1e8 across the two summed axes, so the engine's
|
|
// per-addend narrowing lets 1e8 swallow the running 2 and answers 0,
|
|
// where a fold's narrow-once answer would be 2.
|
|
func TestEinsumReductionKeepsFloat32WithEngine(t *testing.T) {
|
|
a := mustFromFloat32s(t, []float32{1, 1, 1e8, -1e8}, 1, 2, 2)
|
|
out, err := Einsum("ijk->i", a)
|
|
if err != nil {
|
|
t.Fatalf("Einsum: %v", err)
|
|
}
|
|
if out.Dtype() != Float32 {
|
|
t.Fatalf("Einsum answered dtype %s, want float32", out.Dtype())
|
|
}
|
|
if got := out.FloatAt(0); got != 0 {
|
|
t.Fatalf("float32 reduction answered %v, want exactly 0; a narrow-once fold answers 2", got)
|
|
}
|
|
}
|
|
|
|
// TestEinsumReductionFloat32ExactSum pins the plain finite case of the
|
|
// same float32 routing: the engine's per-addend walk answers the exact
|
|
// small-integer sum.
|
|
func TestEinsumReductionFloat32ExactSum(t *testing.T) {
|
|
a := mustFromFloat32s(t, []float32{1, 2, 3, 4}, 1, 2, 2)
|
|
out, err := Einsum("ijk->i", a)
|
|
if err != nil {
|
|
t.Fatalf("Einsum: %v", err)
|
|
}
|
|
if got := out.FloatAt(0); got != 10 {
|
|
t.Fatalf("float32 reduction answered %v, want 10", got)
|
|
}
|
|
}
|
|
|
|
// TestEinsumReductionKeepsComplexWithEngine pins the routing of the
|
|
// single-operand axis reduction for a complex operand: the fold declines
|
|
// it, because the engine's per-slot product seed multiplies the first
|
|
// read by 1+0i, the identity for finite values but not for a purely
|
|
// imaginary infinity, which meets 0 times Inf inside the complex product
|
|
// and turns the real part NaN. A plain addition would carry the
|
|
// infinity through untouched with the real part 2.
|
|
func TestEinsumReductionKeepsComplexWithEngine(t *testing.T) {
|
|
a := mustFromComplexes(t, []complex128{1, 1, complex(0, math.Inf(1)), 0}, 1, 2, 2)
|
|
out, err := Einsum("ijk->i", a)
|
|
if err != nil {
|
|
t.Fatalf("Einsum: %v", err)
|
|
}
|
|
got := out.ComplexAt(0)
|
|
if !math.IsNaN(real(got)) {
|
|
t.Fatalf("real part %v, want NaN from the engine's 1+0i product seed", real(got))
|
|
}
|
|
if imag(got) != math.Inf(1) {
|
|
t.Fatalf("imaginary part %v, want +Inf carried through", imag(got))
|
|
}
|
|
}
|
|
|
|
// TestEinsumReductionComplexExactSum pins the plain finite case of the
|
|
// same complex routing: the seed is the multiplicative identity there,
|
|
// and the exact small-integer sum comes back whole.
|
|
func TestEinsumReductionComplexExactSum(t *testing.T) {
|
|
a := mustFromComplexes(t, []complex128{1 + 1i, 2 + 2i, 3 + 3i, 4 + 4i}, 1, 2, 2)
|
|
out, err := Einsum("ijk->i", a)
|
|
if err != nil {
|
|
t.Fatalf("Einsum: %v", err)
|
|
}
|
|
if got := out.ComplexAt(0); got != 10+10i {
|
|
t.Fatalf("complex reduction answered %v, want (10+10i)", got)
|
|
}
|
|
}
|