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