58 lines
2.0 KiB
Go
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
|
||
|
|
}
|