Files
tensor/internal/core/extrema_test.go
T

188 lines
5.6 KiB
Go
Raw Normal View History

2026-09-03 10:00:00 +02:00
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package core
import (
"math"
"strings"
"testing"
)
func TestMinimumMaximum(t *testing.T) {
a := mustFromInts(t, []int64{1, 5, 3}, 3)
b := mustFromInts(t, []int64{4, 2, 3}, 3)
mn, err := Minimum(a, b)
if err != nil {
t.Fatalf("Minimum: %v", err)
}
if !Equal(mustFromInts(t, []int64{1, 2, 3}, 3), mn) {
t.Fatalf("Minimum: %s", mn)
}
mx, err := Maximum(a, b)
if err != nil {
t.Fatalf("Maximum: %v", err)
}
if !Equal(mustFromInts(t, []int64{4, 5, 3}, 3), mx) {
t.Fatalf("Maximum: %s", mx)
}
// Promotion and NaN propagation.
f := mustFromFloats(t, []float64{1.0, math.NaN()}, 2)
fb := mustFromFloats(t, []float64{0.5, 1.0}, 2)
fmin, _ := Minimum(f, fb)
if v, _ := FloatAt(fmin, 0); v != 0.5 {
t.Fatalf("Minimum promote: %v", v)
}
if v, _ := FloatAt(fmin, 1); !math.IsNaN(v) {
t.Fatalf("Minimum NaN must propagate: %v", v)
}
i := mustFromInts(t, []int64{1}, 1)
mixed, _ := Minimum(i, mustFromFloats(t, []float64{0.5}, 1))
if mixed.Dtype() != Float {
t.Fatalf("Minimum promote dtype: %s", mixed.Dtype())
}
c := mustFromComplexes(t, []complex128{1}, 1)
if _, err := Minimum(c, c); err == nil || !strings.Contains(err.Error(), "no ordering") {
t.Fatalf("Minimum complex: %v", err)
}
if _, err := Maximum(a, mustFromInts(t, []int64{1}, 1)); err == nil || !strings.Contains(err.Error(), "shape mismatch") {
t.Fatalf("Maximum shape: %v", err)
}
}
func TestClip(t *testing.T) {
a := mustFromInts(t, []int64{-5, 3, 99}, 3)
cli, err := ClipI(a, 0, 10)
if err != nil {
t.Fatalf("ClipI: %v", err)
}
if !Equal(mustFromInts(t, []int64{0, 3, 10}, 3), cli) {
t.Fatalf("ClipI: %s", cli)
}
if cli.Dtype() != Int {
t.Fatalf("ClipI keeps int: %s", cli.Dtype())
}
clf, err := ClipF(a, -1.5, 1.5)
if err != nil {
t.Fatalf("ClipF: %v", err)
}
if clf.Dtype() != Float {
t.Fatalf("ClipF dtype: %s", clf.Dtype())
}
if v, _ := FloatAt(clf, 0); v != -1.5 {
t.Fatalf("ClipF lo: %v", v)
}
if v, _ := FloatAt(clf, 2); v != 1.5 {
t.Fatalf("ClipF hi: %v", v)
}
f := mustFromFloats(t, []float64{0.5, 2.5}, 2)
fc, _ := ClipI(f, 1, 2)
if v, _ := FloatAt(fc, 0); v != 1 {
t.Fatalf("ClipI on float: %v", v)
}
if _, err := ClipI(a, 5, 0); err == nil || !strings.Contains(err.Error(), "lo must be at most hi") {
t.Fatalf("ClipI range: %v", err)
}
if _, err := ClipF(a, 2, 1); err == nil || !strings.Contains(err.Error(), "lo must be at most hi") {
t.Fatalf("ClipF range: %v", err)
}
c := mustFromComplexes(t, []complex128{1}, 1)
if _, err := ClipI(c, 0, 1); err == nil || !strings.Contains(err.Error(), "no ordering") {
t.Fatalf("ClipI complex: %v", err)
}
if _, err := ClipF(c, 0, 1); err == nil || !strings.Contains(err.Error(), "no ordering") {
t.Fatalf("ClipF complex: %v", err)
}
}
func TestElementsGeneric(t *testing.T) {
i := mustFromInts(t, []int64{1, 2}, 2)
// Same-type access returns the values unchanged.
ints, err := i.Elements[int64]()
if err != nil || ints[0] != 1 || ints[1] != 2 {
t.Fatalf("Elements[int64]: %v %v", ints, err)
}
// Widening converts along the ladder.
floats, err := i.Elements[float64]()
if err != nil || floats[0] != 1 || floats[1] != 2 {
t.Fatalf("Elements[float64]: %v %v", floats, err)
}
complexes, err := i.Elements[complex128]()
if err != nil || complexes[1] != complex(2, 0) {
t.Fatalf("Elements[complex128]: %v %v", complexes, err)
}
// Float arrays widen to complex; narrowing errors.
f := mustFromFloats(t, []float64{2.5}, 1)
fc, err := f.Elements[complex128]()
if err != nil || fc[0] != complex(2.5, 0) {
t.Fatalf("Elements float to complex: %v %v", fc, err)
}
if _, err := f.Elements[int64](); err == nil || !strings.Contains(err.Error(), "cannot narrow float to int64") {
t.Fatalf("Elements float to int64: %v", err)
}
c := mustFromComplexes(t, []complex128{complex(1, 2)}, 1)
if _, err := c.Elements[float64](); err == nil || !strings.Contains(err.Error(), "cannot narrow complex to float64") {
t.Fatalf("Elements complex to float64: %v", err)
}
cc, err := c.Elements[complex128]()
if err != nil || cc[0] != complex(1, 2) {
t.Fatalf("Elements[complex128] on complex: %v %v", cc, err)
}
// The returned slice is a copy.
vals, _ := i.Elements[int64]()
vals[0] = 99
if v, _ := IntAt(i, 0); v != 1 {
t.Fatalf("Elements must copy: %d", v)
}
// The integer class widens to int64 exactly, the rule IntAt
// carries: a narrow source is an exact widening, never a
// narrowing refusal.
n8, err := FromInt8s([]int8{-128, -1, 0, 1, 127}, 5)
if err != nil {
t.Fatalf("FromInt8s: %v", err)
}
nv, err := n8.Elements[int64]()
if err != nil {
t.Fatalf("Elements[int64] on int8: %v", err)
}
for j, w := range []int64{-128, -1, 0, 1, 127} {
if nv[j] != w {
t.Fatalf("Elements[int64] int8[%d] = %d, want %d", j, nv[j], w)
}
}
bs, err := FromBools([]bool{true, false, true}, 3)
if err != nil {
t.Fatalf("FromBools: %v", err)
}
bi, err := bs.Elements[int64]()
if err != nil || bi[0] != 1 || bi[1] != 0 || bi[2] != 1 {
t.Fatalf("Elements[int64] on bool = %v, %v; want [1 0 1]", bi, err)
}
u32, err := FromUint32s([]uint32{0, 4294967295}, 2)
if err != nil {
t.Fatalf("FromUint32s: %v", err)
}
ui, err := u32.Elements[int64]()
if err != nil || ui[0] != 0 || ui[1] != 4294967295 {
t.Fatalf("Elements[int64] on uint32 = %v, %v; want [0 4294967295] exact", ui, err)
}
// A complex source has no exact int64 image: the genuine
// narrowing keeps its refusal.
if _, err := c.Elements[int64](); err == nil || !strings.Contains(err.Error(), "cannot narrow complex to int64") {
t.Fatalf("Elements complex to int64: %v", err)
}
}