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