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