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

122 lines
3.1 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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
}