122 lines
3.1 KiB
Go
122 lines
3.1 KiB
Go
// 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
|
|||
|
|
}
|