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

60 lines
1.5 KiB
Go

// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package core
import "testing"
func TestOneHotEncoding(t *testing.T) {
codes := &Array{shape: []int{4}, dt: Int, ints: []int64{0, 2, 1, 2}}
hot, err := OneHot(codes, 3)
if err != nil {
t.Fatal(err)
}
if got := hot.Shape(); got[0] != 4 || got[1] != 3 {
t.Fatalf("shape: %v", got)
}
want := [][]float64{{1, 0, 0}, {0, 0, 1}, {0, 1, 0}, {0, 0, 1}}
for i := range 4 {
for j := range 3 {
if g := float64(hot.RawFloat32s()[i*3+j]); g != want[i][j] {
t.Fatalf("hot[%d][%d] = %v, want %v", i, j, g, want[i][j])
}
}
}
// A single-code array keeps its leading dimension and gains the
// class axis at the end.
solo, _ := FromInts([]int64{1}, 1)
hotSolo, err := OneHot(solo, 2)
if err != nil {
t.Fatal(err)
}
if got := hotSolo.Shape(); got[0] != 1 || got[1] != 2 {
t.Fatalf("solo shape: %v", got)
}
}
func TestOneHotErrors(t *testing.T) {
floatCodes, _ := FromFloats([]float64{1}, 1)
if _, err := OneHot(floatCodes, 3); err == nil {
t.Fatal("float codes accepted")
}
intCodes := &Array{shape: []int{2}, dt: Int, ints: []int64{0, 5}}
if _, err := OneHot(intCodes, 3); err == nil {
t.Fatal("out-of-range code accepted")
}
negative := &Array{shape: []int{1}, dt: Int, ints: []int64{-1}}
if _, err := OneHot(negative, 3); err == nil {
t.Fatal("negative code accepted")
}
valid := &Array{shape: []int{1}, dt: Int, ints: []int64{0}}
if _, err := OneHot(valid, 0); err == nil {
t.Fatal("zero classes accepted")
}
}