126 lines
4.4 KiB
Go
126 lines
4.4 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|
// SPDX-License-Identifier: MIT
|
|
|
|
// Package engine hosts the private machinery shared by the tensor
|
|
// packages: the parallel scheduling primitive, the worker-count policy
|
|
// and pooled scratch buffers. It is under internal/: the compiler
|
|
// keeps it invisible outside this module.
|
|
package engine
|
|
|
|
import (
|
|
"runtime"
|
|
"sync"
|
|
)
|
|
|
|
var (
|
|
mu sync.RWMutex
|
|
numWorkers = runtime.NumCPU()
|
|
)
|
|
|
|
// SetNumWorkers sets the number of goroutines the parallel kernels may
|
|
// use and returns the previous value. Values below 1 reset to NumCPU.
|
|
func SetNumWorkers(n int) int {
|
|
mu.Lock()
|
|
defer mu.Unlock()
|
|
prev := numWorkers
|
|
if n < 1 {
|
|
n = runtime.NumCPU()
|
|
}
|
|
numWorkers = n
|
|
return prev
|
|
}
|
|
|
|
// NumWorkers returns the current worker count.
|
|
func NumWorkers() int {
|
|
mu.RLock()
|
|
defer mu.RUnlock()
|
|
return numWorkers
|
|
}
|
|
|
|
// WorkersFor returns the number of goroutines to use for a workload of
|
|
// n independent items, bounded by both the worker count and n.
|
|
func WorkersFor(n int) int { return max(min(NumWorkers(), n), 1) }
|
|
|
|
// Parallel splits the [0, n) index range into chunks and runs fn on
|
|
// each chunk in its own goroutine. Fixed chunk size, no per-item
|
|
// channel traffic; every chunk owns a disjoint slice of the output, so
|
|
// kernels need no locks. A workload the worker policy collapses to a
|
|
// single worker (n of 1, or the worker count pinned to 1) runs inline
|
|
// on the calling goroutine; any other workload spawns one goroutine
|
|
// per worker, so use ParallelMin for a real per-worker floor.
|
|
func Parallel(n int, fn func(start, end int)) { ParallelMin(n, 1, fn) }
|
|
|
|
// ParallelMin splits the [0, n) index range into chunks and runs fn on
|
|
// each chunk in its own goroutine, exactly like Parallel, with one
|
|
// extra constraint: while the per-worker chunk would fall below
|
|
// minPerWorker, fn runs whole as fn(0, n) on the calling goroutine. A
|
|
// worker whose chunk is below the floor costs more to create and
|
|
// schedule than the work it carries, so parallelising that workload
|
|
// only adds latency; the caller's goroutine is already warm and pays
|
|
// nothing to start. Chunk boundaries and the worker choice are
|
|
// computed exactly as Parallel computes them, so minPerWorker values
|
|
// below 2 reproduce Parallel bit for bit.
|
|
func ParallelMin(n, minPerWorker int, fn func(start, end int)) {
|
|
w := WorkersFor(n)
|
|
if w == 1 {
|
|
fn(0, n)
|
|
return
|
|
}
|
|
chunk := (n + w - 1) / w
|
|
if chunk < minPerWorker {
|
|
fn(0, n)
|
|
return
|
|
}
|
|
var wg sync.WaitGroup
|
|
for start := 0; start < n; start += chunk {
|
|
end := min(start+chunk, n)
|
|
wg.Go(func() {
|
|
fn(start, end)
|
|
})
|
|
}
|
|
wg.Wait()
|
|
}
|
|
|
|
var float64Pool = sync.Pool{New: func() any { return make([]float64, 0, 1024) }}
|
|
|
|
// maxPooledFloat64 caps what the scratch pool keeps. A pooled buffer is
|
|
// retained per processor until the next garbage collection, so a kernel
|
|
// that borrows hundreds of megabytes would pin that much memory times
|
|
// the processor count. Buffers above the cap are dropped on return and
|
|
// re-allocated by the next borrower, one allocation per deep chunk;
|
|
// every buffer at or below it still round-trips.
|
|
const maxPooledFloat64 = 1 << 20 // elements, 8 MiB
|
|
|
|
// keepPooled reports whether a returned buffer of the given capacity is
|
|
// worth retaining.
|
|
func keepPooled(capacity int) bool { return capacity <= maxPooledFloat64 }
|
|
|
|
// GetFloat64Buf borrows a float64 buffer of exactly n elements with
|
|
// capacity for at least that many. The buffer may be recycled from an
|
|
// earlier borrower, so it is cleared before it leaves the pool: every
|
|
// slot arrives zero and stays zero until the borrower writes it. That
|
|
// makes accumulation kernels safe by construction: stale sums can
|
|
// never leak into a result, whichever path the buffer took through the
|
|
// pool or the garbage collector.
|
|
func GetFloat64Buf(n int) []float64 {
|
|
b := float64Pool.Get().([]float64)
|
|
if cap(b) < n {
|
|
return make([]float64, n) // freshly allocated: already zero
|
|
}
|
|
b = b[:n]
|
|
clear(b) // pooled buffers come back dirty; never hand that on
|
|
return b
|
|
}
|
|
|
|
// PutFloat64Buf returns a borrowed buffer. The backing array is offered
|
|
// to the next caller, although sync.Pool may drop it at any garbage
|
|
// collection; reuse is opportunistic, never guaranteed. A buffer larger
|
|
// than maxPooledFloat64 is dropped outright so that one deep kernel
|
|
// cannot pin its scratch memory on every processor.
|
|
func PutFloat64Buf(b []float64) {
|
|
if !keepPooled(cap(b)) {
|
|
return
|
|
}
|
|
float64Pool.Put(b[:0])
|
|
}
|