Files
tensor/internal/core/random.go
T

225 lines
7.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
import (
"math"
"math/bits"
)
// Reproducible randomness. The generator runs xoshiro256++
// seeded through splitmix64: implemented here because math/rand/v2 does
// not promise stable output across Go versions, and reproducibility is
// the point of a seed: the same seed yields bit-identical arrays on every
// Go release. The quality suits simulation, not cryptography. The splitmix
// machinery is exported for callers who seed their own streams, and
// Substream hands out the provably distinct members of one seed's stream
// family.
// Generator produces deterministic random arrays from a seed.
type Generator struct {
s [4]uint64
}
// NewGenerator seeds a fresh generator. Any int64 seed is valid.
func NewGenerator(seed int64) *Generator {
return generatorFrom(uint64(seed))
}
// Splitmix64 advances the splitmix64 stream one step: from the given
// state it returns the advanced state and the mixed output. The mixer
// is a bijection on uint64, xor shifts and odd multipliers both
// inverting, which is the property the substream construction rests
// on. Any state is valid, the zero included.
func Splitmix64(state uint64) (uint64, uint64) {
state += 0x9E3779B97F4A7C15
return state, mix64(state)
}
// Substream seeds the index-th member of the stream family one seed
// carries. The start state mixes the index through splitmix64 before
// combining it with the seed and mixing again, and both mixings are
// bijections, so distinct indices give provably distinct initial
// states. That is the property a seed + index stride lacks: its
// streams are one stream read at different offsets, which measures a
// rearrangement where a family of independent streams was wanted.
// Indices start at zero.
func Substream(seed int64, index int) (*Generator, error) {
if index < 0 {
return nil, errf("Substream: the index must be zero or greater, got %d", index)
}
return generatorFrom(mix64(uint64(seed) ^ mix64(uint64(index)))), nil
}
// generatorFrom seeds a generator by walking the splitmix64 stream
// four steps from z.
func generatorFrom(z uint64) *Generator {
g := &Generator{}
for i := range g.s {
var v uint64
z, v = Splitmix64(z)
g.s[i] = v
}
// xoshiro never leaves the all-zero state, so nudge it away.
if g.s == [4]uint64{} {
g.s = [4]uint64{1, 2, 3, 4}
}
return g
}
// mix64 is the splitmix64 finaliser, a bijection on uint64.
func mix64(z uint64) uint64 {
z = (z ^ (z >> 30)) * 0xBF58476D1CE4E5B9
z = (z ^ (z >> 27)) * 0x94D049BB133111EB
return z ^ (z >> 31)
}
// Floats returns n uniform floats in [0, 1) with 53-bit resolution.
func Floats(g *Generator, n int) (*Array, error) {
if n < 0 {
return nil, errf("Floats: n must be zero or greater, got %d", n)
}
floats := make([]float64, n)
for i := range floats {
floats[i] = g.unit()
}
return &Array{shape: []int{n}, dt: Float, floats: floats}, nil
}
// Ints returns n uniform ints in [min, max), drawn without modulo bias
// via Lemire's multiply-shift rejection.
func Ints(g *Generator, n int, min, max int64) (*Array, error) {
if n < 0 {
return nil, errf("Ints: n must be zero or greater, got %d", n)
}
if min >= max {
return nil, errf("Ints: min must be less than max, got %d and %d", min, max)
}
span := uint64(max) - uint64(min)
ints := make([]int64, n)
for i := range ints {
ints[i] = min + int64(g.bounded(span))
}
return &Array{shape: []int{n}, dt: Int, ints: ints}, nil
}
// bounded returns an unbiased uniform value in [0, n) for n > 0, via
// Lemire's multiply-shift rejection. The 128-bit product x*n splits into
// the candidate hi and residue lo; the rejected band has exactly
// t = 2^64 mod n residues, so a draw is accepted once lo >= t (lo >= n
// already implies that, deferring the modulo to the rare retry path).
func (g *Generator) bounded(n uint64) uint64 {
hi, lo := bits.Mul64(g.next(), n)
if lo < n {
t := -n % n // 2^64 mod n
for lo < t {
hi, lo = bits.Mul64(g.next(), n)
}
}
return hi
}
// Normal returns n Gaussian draws of the given mean and standard
// deviation, paired Box-Muller transforms of the xoshiro stream,
// bit-stable across Go releases like every other draw. A
// negative or NaN std is an error.
func Normal(g *Generator, n int, mean, std float64) (*Array, error) {
if n < 0 {
return nil, errf("Normal: n must be zero or greater, got %d", n)
}
if std < 0 || math.IsNaN(std) {
return nil, errf("Normal: std must be zero or greater, got %v", std)
}
floats := make([]float64, n)
for i := range n {
floats[i] = mean + std*g.normalUnit()
}
return &Array{shape: []int{n}, dt: Float, floats: floats}, nil
}
// normalUnit draws one standard normal value via the polar Box-Muller
// method: rejection until a point lands in the unit disk, then the
// spare-free transform.
func (g *Generator) normalUnit() float64 {
for {
u := 2*g.unit() - 1
v := 2*g.unit() - 1
if s := u*u + v*v; s < 1 && s > 0 {
return u * math.Sqrt(-2*math.Log(s)/s)
}
}
}
// unit draws one uniform float in [0, 1) with 53-bit resolution.
func (g *Generator) unit() float64 {
return float64(g.next()>>11) / (1 << 53)
}
// Permutation returns a uniform permutation of 0..n-1 (Fisher-Yates
// over unbiased bounded draws). Renamed from `Permute` to avoid
// confusion with `TransposeAxes` (which used to be called Permute
// before it became the axis-permutation variant).
func Permutation(g *Generator, n int) (*Array, error) {
if n < 0 {
return nil, errf("Permutation: n must be zero or greater, got %d", n)
}
ints := make([]int64, n)
for i := range ints {
ints[i] = int64(i)
}
for i := n - 1; i > 0; i-- {
j := g.bounded(uint64(i) + 1)
ints[i], ints[j] = ints[j], ints[i]
}
return &Array{shape: []int{n}, dt: Int, ints: ints}, nil
}
// Shuffle returns a shuffled copy of a; the receiver itself is never
// touched.
func Shuffle(g *Generator, a *Array) *Array {
perm, _ := Permutation(g, a.Len())
out := &Array{shape: a.Shape(), dt: a.dt}
out.alloc(a.Len())
for i := range a.Len() {
out.setFrom(i, a, int(perm.ints[i]))
}
return out
}
// next advances xoshiro256++ and returns the next 64-bit value.
func (g *Generator) next() uint64 {
result := bits.RotateLeft64(g.s[0]+g.s[3], 23) + g.s[0]
t := g.s[1] << 17
g.s[2] ^= g.s[0]
g.s[3] ^= g.s[1]
g.s[1] ^= g.s[2]
g.s[0] ^= g.s[3]
g.s[2] ^= t
g.s[3] = bits.RotateLeft64(g.s[3], 45)
return result
}
// Float32s returns n uniform float32 values in [0, 1) with 24-bit
// resolution.
func Float32s(g *Generator, n int) (*Array, error) {
if n < 0 {
return nil, errf("Float32s: n must be zero or greater, got %d", n)
}
floats32 := make([]float32, n)
for i := range floats32 {
floats32[i] = float32(g.next()>>40) / (1 << 24)
}
return &Array{shape: []int{n}, dt: Float32, floats32: floats32}, nil
}
// Unit draws one uniform float in [0, 1) with 53-bit resolution.
func (g *Generator) Unit() float64 { return g.unit() }
// NormalUnit draws one standard normal value via the polar
// Box-Muller method.
func (g *Generator) NormalUnit() float64 { return g.normalUnit() }
// Next advances the generator and returns the raw 64-bit value.
func (g *Generator) Next() uint64 { return g.next() }