220 lines
7.8 KiB
Go
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)
|
|
}
|
|
}
|
|
})
|
|
}
|