feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,59 @@
|
||||
// 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")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user