Files
tensor/internal/core/init.go
T

122 lines
3.1 KiB
Go
Raw Normal View History

2026-09-03 10:00:00 +02:00
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package core
// TruncatedNormal returns an array of the given shape with values
// drawn from a normal distribution truncated to ±2σ. Rejection sampling
// redraws until a value lands inside the window. All draws share the
// Generator for reproducibility; the dtype is float32. An invalid
// shape yields a nil array, like the other no-error constructors, and
// a std that is zero, negative or NaN is degenerate: the answer is the
// all-zero array, the escape from a retry loop that negative and NaN
// values could never leave.
func TruncatedNormal(g *Generator, shape []int, mean, std float64) *Array {
n, sh, err := checkedDims(shape)
if err != nil {
return nil
}
if !(std > 0) {
// std <= 0 or NaN: the window is a point, inverted or undefined,
// and the negative and NaN cases could never leave the retry loop
// below, so the degenerate fallback returns the all-zero array.
out := &Array{shape: sh, dt: Float32}
out.alloc(n)
return out
}
out := make([]float32, n)
twoStd := 2.0 * std
for i := range n {
// The scalar draw consumes the generator exactly as one
// Normal(g, 1, 0, std) call would (std is validated above, so
// that call cannot error), without its per-element allocation.
for {
v := std * g.normalUnit()
if v >= -twoStd && v <= twoStd {
out[i] = float32(v + mean)
break
}
}
}
return &Array{shape: sh, dt: Float32, floats32: out}
}
// ZerosLike and OnesLike return a zero- or one-filled array with the
// same shape and dtype as a. The result is a copy: mutations on it
// do not affect a.
func ZerosLike(a *Array) *Array {
return fullLike(a, 0)
}
func OnesLike(a *Array) *Array {
return fullLike(a, 1)
}
// fullLike creates a new array with a's shape and dtype, filled with v.
// The narrow integer widths and bool take the same implicit-store cast
// every other fill path carries: Go's conversion through int64, and
// v != 0 for bool.
func fullLike(a *Array, v float64) *Array {
out := &Array{shape: a.Shape(), dt: a.dt}
out.alloc(a.Len())
switch a.dt {
case Int:
for i := range out.ints {
out.ints[i] = int64(v)
}
case Bool:
bv := v != 0
for i := range out.bools {
out.bools[i] = bv
}
case Int8:
w := int8(int64(v))
for i := range out.i8s {
out.i8s[i] = w
}
case Uint8:
w := uint8(int64(v))
for i := range out.u8s {
out.u8s[i] = w
}
case Int16:
w := int16(int64(v))
for i := range out.i16s {
out.i16s[i] = w
}
case Uint16:
w := uint16(int64(v))
for i := range out.u16s {
out.u16s[i] = w
}
case Int32:
w := int32(int64(v))
for i := range out.i32s {
out.i32s[i] = w
}
case Uint32:
w := uint32(int64(v))
for i := range out.u32s {
out.u32s[i] = w
}
case Float16:
hv := HalfFromFloat64(v)
for i := range out.halves {
out.halves[i] = hv
}
case Float32:
for i := range out.floats32 {
out.floats32[i] = float32(v)
}
case Float:
for i := range out.floats {
out.floats[i] = v
}
default:
for i := range out.complexes {
out.complexes[i] = complex(v, 0)
}
}
return out
}