Files
tensor/internal/core/elementwise_view_test.go
T
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

220 lines
7.8 KiB
Go

// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package core
import "testing"
// TestElementwiseOpTakesTheStridedFallback pins the contiguity guard of
// elementwiseOp: operands carrying a stride table keep out of the raw
// payload kernels and take the accessor walk instead. No public
// constructor sets strides, so the guard is driven here directly: the
// view arrays carry their own stride tables and the unexported entry
// point is called in package. The result must read the strided elements
// in logical order and come back dense, and neither operand's storage
// may move.
func TestElementwiseOpTakesTheStridedFallback(t *testing.T) {
t.Run("same dtype, non-canonical strides", func(t *testing.T) {
// The strides gather rows 0 and 2 of a (3, 2) storage, so a
// payload-order read answers 1, 2, 3, 4 while the logical one
// answers 1, 2, 4, 5: the same-width raw-payload kernel must
// stay out of reach behind the dense gate, not just behind the
// dispatcher's stride disjunct.
a := &Array{shape: []int{2, 2}, dt: Float, floats: []float64{1, 2, 3, 4, 5, 6}, strides: []int{3, 1}}
b := &Array{shape: []int{2, 2}, dt: Float, floats: []float64{10, 20, 30, 40}}
out, err := elementwiseOp(a, b, "Add", pairAdd)
if err != nil {
t.Fatalf("elementwiseOp: %v", err)
}
want := []float64{11, 22, 34, 45}
for i, w := range want {
if got := out.FloatAt(i); got != w {
t.Fatalf("element %d = %v, want %v", i, got, w)
}
}
if out.Strided() {
t.Fatal("the result came back strided, want a dense payload")
}
})
t.Run("mixed dtypes, rebased window", func(t *testing.T) {
// The int view gathers every second row of a (3, 2) storage, so
// its logical elements are 1, 2, 4, 5. Being the lower operand of
// the mixed pair, it is read through the accessor walk, which
// resolves its stride table.
ai := []int64{1, 2, 3, 4, 5, 6}
a := &Array{shape: []int{2, 2}, dt: Int, ints: ai, strides: []int{3, 1}}
b := mustFloats(t, []float64{10, 20, 30, 40}, 2, 2)
out, err := elementwiseOp(a, b, "Add", pairAdd)
if err != nil {
t.Fatalf("elementwiseOp: %v", err)
}
if out.Dtype() != Float {
t.Fatalf("elementwiseOp answered dtype %s, want float", out.Dtype())
}
want := []float64{11, 22, 34, 45}
for i, w := range want {
if got := out.FloatAt(i); got != w {
t.Fatalf("element %d = %v, want %v", i, got, w)
}
}
if out.Strided() {
t.Fatal("the result came back strided, want a dense payload")
}
for i, v := range ai {
if v != int64(i+1) {
t.Fatalf("the operand's storage moved at %d", i)
}
}
})
}
// TestElementwiseStridedOperandsReadLogically pins the gate that keeps
// the accessor walk honest inside elementwise and elementwiseDiv: every
// raw-payload branch is reserved to dense operands, so a strided view
// reads in its logical order on every entry point that reaches the
// fallback, Div and the extrema family included. The view gathers rows
// 0 and 2 of a (3, 2) storage, so its logical elements are 1, 2, 4, 5
// while a payload-order read answers 1, 2, 3, 4.
func TestElementwiseStridedOperandsReadLogically(t *testing.T) {
storage := []float64{1, 2, 3, 4, 5, 6}
view := &Array{shape: []int{2, 2}, dt: Float, floats: storage, strides: []int{3, 1}}
denseB := &Array{shape: []int{2, 2}, dt: Float, floats: []float64{10, 20, 30, 40}}
want := []float64{11, 22, 34, 45}
t.Run("add float64", func(t *testing.T) {
out, err := Add(view, denseB)
if err != nil {
t.Fatalf("Add: %v", err)
}
for i, w := range want {
if got := out.floats[i]; got != w {
t.Fatalf("element %d = %v, want %v", i, got, w)
}
}
})
t.Run("minimum float64", func(t *testing.T) {
out, err := Minimum(view, denseB)
if err != nil {
t.Fatalf("Minimum: %v", err)
}
for i, w := range []float64{1, 2, 4, 5} {
if got := out.floats[i]; got != w {
t.Fatalf("element %d = %v, want %v", i, got, w)
}
}
})
t.Run("div float32", func(t *testing.T) {
a := &Array{shape: []int{2, 2}, dt: Float32, floats32: []float32{1, 2, 3, 4, 5, 6}, strides: []int{3, 1}}
b := &Array{shape: []int{2, 2}, dt: Float32, floats32: []float32{4, 8, 16, 40}}
out, err := Div(a, b)
if err != nil {
t.Fatalf("Div: %v", err)
}
for i, w := range []float32{0.25, 0.25, 0.25, 0.125} {
if got := out.floats32[i]; got != w {
t.Fatalf("element %d = %v, want %v", i, got, w)
}
}
})
t.Run("add float16", func(t *testing.T) {
halves := make([]uint16, len(storage))
for i, v := range storage {
halves[i] = HalfFromFloat64(v)
}
a := &Array{shape: []int{2, 2}, dt: Float16, halves: halves, strides: []int{3, 1}}
bh := []uint16{HalfFromFloat64(10), HalfFromFloat64(20), HalfFromFloat64(30), HalfFromFloat64(40)}
b := &Array{shape: []int{2, 2}, dt: Float16, halves: bh}
out, err := Add(a, b)
if err != nil {
t.Fatalf("Add: %v", err)
}
for i, w := range want {
if got := HalfToFloat64(out.halves[i]); got != w {
t.Fatalf("element %d = %v, want %v", i, got, w)
}
}
})
t.Run("add int64", func(t *testing.T) {
ints := []int64{1, 2, 3, 4, 5, 6}
a := &Array{shape: []int{2, 2}, dt: Int, ints: ints, strides: []int{3, 1}}
b := &Array{shape: []int{2, 2}, dt: Int, ints: []int64{10, 20, 30, 40}}
out, err := Add(a, b)
if err != nil {
t.Fatalf("Add: %v", err)
}
for i, w := range []int64{11, 22, 34, 45} {
if got := out.ints[i]; got != w {
t.Fatalf("element %d = %v, want %v", i, got, w)
}
}
for i, v := range ints {
if v != int64(i+1) {
t.Fatalf("the operand's storage moved at %d", i)
}
}
})
t.Run("div float64", func(t *testing.T) {
a := &Array{shape: []int{2, 2}, dt: Float, floats: []float64{1, 2, 3, 4, 5, 6}, strides: []int{3, 1}}
b := &Array{shape: []int{2, 2}, dt: Float, floats: []float64{4, 8, 16, 40}}
out, err := Div(a, b)
if err != nil {
t.Fatalf("Div: %v", err)
}
for i, w := range []float64{0.25, 0.25, 0.25, 0.125} {
if got := out.floats[i]; got != w {
t.Fatalf("element %d = %v, want %v", i, got, w)
}
}
})
t.Run("div float16", func(t *testing.T) {
halves := make([]uint16, 6)
for i, v := range []float64{1, 2, 3, 4, 5, 6} {
halves[i] = HalfFromFloat64(v)
}
a := &Array{shape: []int{2, 2}, dt: Float16, halves: halves, strides: []int{3, 1}}
bh := []uint16{HalfFromFloat64(4), HalfFromFloat64(8), HalfFromFloat64(16), HalfFromFloat64(40)}
b := &Array{shape: []int{2, 2}, dt: Float16, halves: bh}
out, err := Div(a, b)
if err != nil {
t.Fatalf("Div: %v", err)
}
for i, w := range []float64{0.25, 0.25, 0.25, 0.125} {
if got := HalfToFloat64(out.halves[i]); got != w {
t.Fatalf("element %d = %v, want %v", i, got, w)
}
}
})
t.Run("add complex128", func(t *testing.T) {
complexes := []complex128{1 + 1i, 2 + 2i, 3 + 3i, 4 + 4i, 5 + 5i, 6 + 6i}
a := &Array{shape: []int{2, 2}, dt: Complex, complexes: complexes, strides: []int{3, 1}}
b := &Array{shape: []int{2, 2}, dt: Complex, complexes: []complex128{10, 20, 30, 40}}
out, err := Add(a, b)
if err != nil {
t.Fatalf("Add: %v", err)
}
for i, w := range []complex128{11 + 1i, 22 + 2i, 34 + 4i, 45 + 5i} {
if got := out.complexes[i]; got != w {
t.Fatalf("element %d = %v, want %v", i, got, w)
}
}
for i, v := range complexes {
if v != complex(float64(i+1), float64(i+1)) {
t.Fatalf("the operand's storage moved at %d", i)
}
}
})
t.Run("div complex128", func(t *testing.T) {
a := &Array{shape: []int{2, 2}, dt: Complex, complexes: []complex128{1 + 1i, 2 + 2i, 3 + 3i, 4 + 4i, 5 + 5i, 6 + 6i}, strides: []int{3, 1}}
b := &Array{shape: []int{2, 2}, dt: Complex, complexes: []complex128{4, 8, 16, 40}}
out, err := Div(a, b)
if err != nil {
t.Fatalf("Div: %v", err)
}
for i, w := range []complex128{0.25 + 0.25i, 0.25 + 0.25i, 0.25 + 0.25i, 0.125 + 0.125i} {
if got := out.complexes[i]; got != w {
t.Fatalf("element %d = %v, want %v", i, got, w)
}
}
})
}