60 lines
1.5 KiB
Go
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")
|
||
|
|
}
|
||
|
|
}
|