Files
tensor/internal/core/concat_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

233 lines
6.7 KiB
Go

// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package core
import (
"slices"
"strings"
"testing"
)
func TestConcat(t *testing.T) {
a := mustFromInts(t, []int64{1, 2, 3, 4}, 2, 2)
b := mustFromInts(t, []int64{5, 6, 7, 8}, 2, 2)
// Along columns: rows grow wider.
wide, err := Concat(a, b, 1)
if err != nil {
t.Fatalf("Concat: %v", err)
}
want := mustFromInts(t, []int64{1, 2, 5, 6, 3, 4, 7, 8}, 2, 4)
if !Equal(want, wide) {
t.Fatalf("Concat dim 1: %s", wide)
}
// Along rows: the array grows taller.
tall, err := Concat(a, b, 0)
if err != nil {
t.Fatalf("Concat dim 0: %v", err)
}
wantTall := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6, 7, 8}, 4, 2)
if !Equal(wantTall, tall) {
t.Fatalf("Concat dim 0: %s", tall)
}
// Promotion across operands.
f := mustFromFloats(t, []float64{0.5, 0.5}, 1, 2)
mixed, err := Concat(a, f, 0)
if err != nil || mixed.Dtype() != Float {
t.Fatalf("Concat promote: %s %v", mixed, err)
}
if v, _ := FloatAt(mixed, 2, 0); v != 0.5 {
t.Fatalf("Concat promote value: %v", v)
}
if _, err := Concat(a, b, 2); err == nil || !strings.Contains(err.Error(), "out of range") {
t.Fatalf("Concat dim range: %v", err)
}
other := mustFromInts(t, []int64{1, 2, 3}, 1, 3)
if _, err := Concat(a, other, 0); err == nil || !strings.Contains(err.Error(), "disagree outside dimension") {
t.Fatalf("Concat shape: %v", err)
}
v := mustFromInts(t, []int64{1}, 1)
if _, err := Concat(a, v, 0); err == nil || !strings.Contains(err.Error(), "ranks differ") {
t.Fatalf("Concat rank: %v", err)
}
}
func TestStack(t *testing.T) {
a := mustFromInts(t, []int64{1, 2, 3, 4}, 2, 2)
b := mustFromInts(t, []int64{5, 6, 7, 8}, 2, 2)
s, err := Stack(a, b, 0)
if err != nil {
t.Fatalf("Stack: %v", err)
}
if shape := s.Shape(); shape[0] != 2 || shape[1] != 2 || shape[2] != 2 {
t.Fatalf("Stack shape: %v", shape)
}
if v, _ := IntAt(s, 0, 0, 0); v != 1 {
t.Fatalf("Stack[0,0,0]: %d", v)
}
if v, _ := IntAt(s, 1, 1, 1); v != 8 {
t.Fatalf("Stack[1,1,1]: %d", v)
}
// Inserting a middle axis keeps the source order.
mid, err := Stack(a, b, 1)
if err != nil {
t.Fatalf("Stack mid: %v", err)
}
if shape := mid.Shape(); shape[0] != 2 || shape[1] != 2 || shape[2] != 2 {
t.Fatalf("Stack mid shape: %v", shape)
}
// mid[0][1][1] is b[0][1]; the inserted axis selects the operand.
if v, _ := IntAt(mid, 0, 1, 1); v != 6 {
t.Fatalf("Stack mid value: %d", v)
}
// Promotion and errors.
f := mustFromFloats(t, []float64{9, 9, 9, 9}, 2, 2)
sf, err := Stack(a, f, 2)
if err != nil || sf.Dtype() != Float {
t.Fatalf("Stack promote: %s %v", sf, err)
}
if _, err := Stack(a, mustFromInts(t, []int64{1}, 1), 0); err == nil || !strings.Contains(err.Error(), "must be identical") {
t.Fatalf("Stack shape: %v", err)
}
if _, err := Stack(a, b, 3); err == nil || !strings.Contains(err.Error(), "out of range") {
t.Fatalf("Stack dim: %v", err)
}
}
// TestConcatStackUnequalParts pins the run copies of Concat and Stack
// when the operands differ in length along the joined axis and several
// outer positions precede it: the second operand's run starts at a
// different source offset per position and its length is its own.
func TestConcatStackUnequalParts(t *testing.T) {
av := make([]int64, 12)
for i := range av {
av[i] = int64(i + 1)
}
bv := make([]int64, 18)
for i := range bv {
bv[i] = int64(101 + i)
}
a := mustFromInts(t, av, 2, 2, 3)
b := mustFromInts(t, bv, 2, 3, 3)
c, err := Concat(a, b, 1)
if err != nil {
t.Fatalf("Concat unequal dim 1: %v", err)
}
if sh := c.Shape(); len(sh) != 3 || sh[0] != 2 || sh[1] != 5 || sh[2] != 3 {
t.Fatalf("Concat unequal dim 1 shape: %v, want [2 5 3]", sh)
}
for i := range 2 {
for j := range 5 {
for k := range 3 {
var want int64
if j < 2 {
want, _ = IntAt(a, i, j, k)
} else {
want, _ = IntAt(b, i, j-2, k)
}
if got, _ := IntAt(c, i, j, k); got != want {
t.Fatalf("Concat unequal dim 1 [%d %d %d] = %d, want %d", i, j, k, got, want)
}
}
}
}
// The last axis, where every leading position is its own head: two
// elements of one operand and one of the other.
av2 := make([]int64, 12)
for i := range av2 {
av2[i] = int64(i + 1)
}
a2 := mustFromInts(t, av2, 2, 3, 2)
b2 := mustFromInts(t, []int64{201, 202, 203, 204, 205, 206}, 2, 3, 1)
c2, err := Concat(a2, b2, 2)
if err != nil {
t.Fatalf("Concat unequal dim 2: %v", err)
}
if sh := c2.Shape(); len(sh) != 3 || sh[0] != 2 || sh[1] != 3 || sh[2] != 3 {
t.Fatalf("Concat unequal dim 2 shape: %v, want [2 3 3]", sh)
}
for i := range 2 {
for j := range 3 {
for k := range 3 {
var want int64
if k < 2 {
want, _ = IntAt(a2, i, j, k)
} else {
want, _ = IntAt(b2, i, j, k-2)
}
if got, _ := IntAt(c2, i, j, k); got != want {
t.Fatalf("Concat unequal dim 2 [%d %d %d] = %d, want %d", i, j, k, got, want)
}
}
}
}
// Stack inserts an axis of two: each outer position contributes one
// run of each operand.
s1 := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6}, 2, 3)
s2 := mustFromInts(t, []int64{11, 12, 13, 14, 15, 16}, 2, 3)
st, err := Stack(s1, s2, 1)
if err != nil {
t.Fatalf("Stack: %v", err)
}
if sh := st.Shape(); len(sh) != 3 || sh[0] != 2 || sh[1] != 2 || sh[2] != 3 {
t.Fatalf("Stack shape: %v, want [2 2 3]", sh)
}
for i := range 2 {
for k := range 3 {
if got, _ := IntAt(st, i, 0, k); got != int64(i*3+k+1) {
t.Fatalf("Stack [%d 0 %d] = %d, want %d", i, k, got, i*3+k+1)
}
if got, _ := IntAt(st, i, 1, k); got != int64(11+i*3+k) {
t.Fatalf("Stack [%d 1 %d] = %d, want %d", i, k, got, 11+i*3+k)
}
}
}
}
// TestConcatNarrowDtypes pins that the narrow payloads ride the
// converted walk until the run copy dispatches them: same-dtype joins
// keep their dtype and mixed joins promote through the containment
// table.
func TestConcatNarrowDtypes(t *testing.T) {
a, err := FromInt8s([]int8{1, 2}, 2)
if err != nil {
t.Fatal(err)
}
b, err := FromInt8s([]int8{3}, 1)
if err != nil {
t.Fatal(err)
}
joined, err := Concat(a, b, 0)
if err != nil {
t.Fatalf("Concat int8: %v", err)
}
if joined.Dtype() != Int8 {
t.Fatalf("Concat int8 answered %s, want int8", joined.Dtype())
}
if want := []int8{1, 2, 3}; !slices.Equal(joined.RawInt8s(), want) {
t.Fatalf("Concat int8 = %v, want %v", joined.RawInt8s(), want)
}
// Mixed signedness promotes: int8 with uint8 answers int16.
u, err := FromUint8s([]uint8{200}, 1)
if err != nil {
t.Fatal(err)
}
mix, err := Concat(a, u, 0)
if err != nil {
t.Fatalf("Concat int8 with uint8: %v", err)
}
if mix.Dtype() != Int16 || !slices.Equal(mix.RawInt16s(), []int16{1, 2, 200}) {
t.Fatalf("Concat int8 with uint8 = %s %v", mix.Dtype(), mix.RawInt16s())
}
}