Files
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

58 lines
2.0 KiB
Go

// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package core
import "sourcedock.dev/petrbalvin/tensor/internal/base"
import (
"math"
"sourcedock.dev/petrbalvin/tensor/internal/engine"
)
// One-hot encoding. The sole producer here is the lookup pathway:
// a differentiable row-gather collapses to one-hot·weights, which the
// existing matrix-product backward already differentiates exactly, so
// embedding gradients need no dedicated kernel.
// OneHot expands integer class codes into indicator vectors: every code
// becomes a length-classes row of zeros with a single one at the code's
// position, appended as a NEW LAST AXIS, so shape (…) turns into
// (…, classes). Codes must be an int array and each value inside
// [0, classes). The result is float32: indicators are exactly
// representable there and downstream weight matrices stay in the fast
// float32 lane.
func OneHot(codes *Array, classes int) (*Array, error) {
if codes.dt != Int {
return nil, errf("OneHot: codes must be an int array, got %s", codes.dt)
}
if classes < 1 {
return nil, errf("OneHot: classes must be positive, got %d", classes)
}
n := codes.Len()
for i := range n {
if c := codes.ints[i]; c < 0 || int(c) >= classes {
return nil, errf("OneHot: code %d out of range [0, %d) at position %d", c, classes, i)
}
}
// Bound the row length before multiplying: a wrapped product must be
// an error, not a negative allocation.
if classes != 0 && n > math.MaxInt/classes {
return nil, errf("OneHot: %d rows of width %d hold more elements than fit in an index", n, classes)
}
total, shape, terr := checkedDims(append(codes.Shape(), classes))
if terr != nil {
return nil, base.WrapErr("OneHot", terr)
}
out := &Array{shape: shape, dt: Float32}
out.floats32 = make([]float32, total)
engine.Parallel(n, func(rs, re int) {
for i := rs; i < re; i++ {
out.floats32[i*classes+int(codes.ints[i])] = 1
}
})
return out, nil
}