// Copyright (c) 2026 Petr Balvín (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 }