feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,125 @@
|
||||
// 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])
|
||||
}
|
||||
@@ -0,0 +1,25 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package engine
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// BenchmarkPoolRoundTrip measures the full borrow/return cycle at the
|
||||
// sizes the kernels actually request, including the clear-on-borrow
|
||||
// cost the pool contract pays. Run before and after touching the pool.
|
||||
func BenchmarkPoolRoundTrip(b *testing.B) {
|
||||
for _, n := range []int{64, 4096, 262144} {
|
||||
b.Run(fmt.Sprintf("n=%d", n), func(b *testing.B) {
|
||||
b.ReportAllocs()
|
||||
for b.Loop() {
|
||||
buf := GetFloat64Buf(n)
|
||||
buf[0] = 1 // touch one slot: prove the buffer is writable
|
||||
PutFloat64Buf(buf)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,171 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package engine
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// TestParallelCoversEveryIndexExactlyOnce is the core invariant: the
|
||||
// chunks partition [0, n), whatever the worker count does to their
|
||||
// boundaries.
|
||||
func TestParallelCoversEveryIndexExactlyOnce(t *testing.T) {
|
||||
for _, n := range []int{0, 1, 2, 7, 33, 100, 1024} {
|
||||
touched := make([]int, n)
|
||||
// The chunk check records violations instead of calling
|
||||
// Fatalf from inside the worker goroutines: FailNow is defined
|
||||
// for the test's own goroutine only. The append sits under a
|
||||
// mutex because the chunks run concurrently.
|
||||
var mu sync.Mutex
|
||||
illegal := make([][3]int, 0, 4)
|
||||
Parallel(n, func(start, end int) {
|
||||
if start < 0 || end > n || start > end {
|
||||
mu.Lock()
|
||||
illegal = append(illegal, [3]int{n, start, end})
|
||||
mu.Unlock()
|
||||
}
|
||||
for i := start; i < end; i++ {
|
||||
touched[i]++
|
||||
}
|
||||
})
|
||||
if len(illegal) > 0 {
|
||||
t.Fatalf("illegal chunks: %v", illegal)
|
||||
}
|
||||
for i, c := range touched {
|
||||
if c != 1 && n > 0 {
|
||||
t.Fatalf("n=%d: index %d visited %d times", n, i, c)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestParallelSmallWorkloadRunsInline pins the small-workload rule:
|
||||
// with a single effective worker the callback runs on the caller's
|
||||
// goroutine before Parallel returns: no goroutine churn for tiny
|
||||
// kernels. A panicked chunk therefore crashes this test instead of
|
||||
// hiding behind the WaitGroup.
|
||||
func TestParallelSmallWorkloadRunsInline(t *testing.T) {
|
||||
prev := SetNumWorkers(1)
|
||||
defer SetNumWorkers(prev)
|
||||
|
||||
called := false
|
||||
Parallel(4, func(start, end int) {
|
||||
called = true
|
||||
if start != 0 || end != 4 {
|
||||
t.Fatalf("single-worker chunk [%d, %d), want [0, 4)", start, end)
|
||||
}
|
||||
})
|
||||
if !called {
|
||||
t.Fatal("callback never ran")
|
||||
}
|
||||
}
|
||||
|
||||
// TestWorkersForBounds checks both ceilings and the floor.
|
||||
func TestWorkersForBounds(t *testing.T) {
|
||||
prev := SetNumWorkers(8)
|
||||
defer SetNumWorkers(prev)
|
||||
|
||||
for _, tc := range []struct{ n, want int }{
|
||||
{0, 1}, {1, 1}, {3, 3}, {8, 8}, {500, 8},
|
||||
} {
|
||||
if got := WorkersFor(tc.n); got != tc.want {
|
||||
t.Errorf("WorkersFor(%d) = %d, want %d", tc.n, got, tc.want)
|
||||
}
|
||||
}
|
||||
|
||||
if got := SetNumWorkers(0); got != 8 {
|
||||
t.Errorf("SetNumWorkers(0) reported previous %d, want 8", got)
|
||||
}
|
||||
if NumWorkers() < 1 {
|
||||
t.Error("reset landed on an unusable worker count")
|
||||
}
|
||||
}
|
||||
|
||||
// TestFloat64PoolRoundTripKeepsLengthAndCapacity pins the borrow
|
||||
// contract: every GetFloat64Buf returns exactly the requested length
|
||||
// with capacity for at least that many, buffers survive a Put/Get round
|
||||
// trip as usable memory, and a buffer handed back dirty arrives cleared.
|
||||
// The pool enforces the zero-on-borrow guarantee itself, so no kernel
|
||||
// can leak an earlier borrower's sums into its result (the regression
|
||||
// TestMatMul2DFloat32ScratchCleaned pins end-to-end in the tensor
|
||||
// package).
|
||||
func TestFloat64PoolRoundTripKeepsLengthAndCapacity(t *testing.T) {
|
||||
buf := GetFloat64Buf(64)
|
||||
if len(buf) != 64 {
|
||||
t.Fatalf("borrowed len %d, want 64", len(buf))
|
||||
}
|
||||
if cap(buf) < 64 {
|
||||
t.Fatalf("borrowed cap %d, want at least 64", cap(buf))
|
||||
}
|
||||
for i := range buf {
|
||||
buf[i] = float64(i) // fill the whole window: prove it is writable
|
||||
}
|
||||
PutFloat64Buf(buf)
|
||||
|
||||
again := GetFloat64Buf(32)
|
||||
if len(again) != 32 {
|
||||
t.Fatalf("re-borrowed len %d, want 32", len(again))
|
||||
}
|
||||
if cap(again) < 32 {
|
||||
t.Fatalf("re-borrowed capacity %d, want at least 32", cap(again))
|
||||
}
|
||||
// The dirty residue the test just put back must never surface: the
|
||||
// pool clears on borrow, so the window arrives all zeros. (sync.Pool
|
||||
// may also drop the buffer at any GC, in which case a fresh (and
|
||||
// therefore zeroed) allocation takes its place; the guarantee holds
|
||||
// on both paths.)
|
||||
for i := range again {
|
||||
if again[i] != 0 {
|
||||
t.Fatalf("slot %d = %v on arrival, want 0", i, again[i])
|
||||
}
|
||||
again[i] = float64(i)
|
||||
if again[i] != float64(i) {
|
||||
t.Fatalf("slot %d = %v after write, want %v", i, again[i], float64(i))
|
||||
}
|
||||
}
|
||||
PutFloat64Buf(again)
|
||||
}
|
||||
|
||||
// TestFloat64PoolGrowsForLargerBorrow checks the grow path returns a
|
||||
// slice of exactly the requested length even when the pooled buffer
|
||||
// must be reallocated.
|
||||
func TestFloat64PoolGrowsForLargerBorrow(t *testing.T) {
|
||||
small := GetFloat64Buf(4)
|
||||
PutFloat64Buf(small)
|
||||
|
||||
big := GetFloat64Buf(4096)
|
||||
if len(big) != 4096 {
|
||||
t.Fatalf("grown len %d, want 4096", len(big))
|
||||
}
|
||||
big[4095] = 1 // writable end to end
|
||||
PutFloat64Buf(big)
|
||||
}
|
||||
|
||||
// TestFloat64PoolRetentionCap pins the size rule: the cap itself
|
||||
// round-trips, anything above it is dropped rather than retained per
|
||||
// processor, and the drop path leaves the pool usable.
|
||||
func TestFloat64PoolRetentionCap(t *testing.T) {
|
||||
if !keepPooled(maxPooledFloat64) {
|
||||
t.Fatalf("a buffer of the cap (%d) must be retained", maxPooledFloat64)
|
||||
}
|
||||
if keepPooled(maxPooledFloat64 + 1) {
|
||||
t.Fatalf("a buffer above the cap (%d) must be dropped", maxPooledFloat64+1)
|
||||
}
|
||||
|
||||
oversized := GetFloat64Buf(4 * maxPooledFloat64)
|
||||
oversized[len(oversized)-1] = 1
|
||||
PutFloat64Buf(oversized) // dropped: must not corrupt the pool
|
||||
|
||||
next := GetFloat64Buf(16)
|
||||
if len(next) != 16 {
|
||||
t.Fatalf("borrowed after a dropped buffer: len %d, want 16", len(next))
|
||||
}
|
||||
for i := range next {
|
||||
if next[i] != 0 {
|
||||
t.Fatalf("slot %d = %v after a dropped buffer, want 0", i, next[i])
|
||||
}
|
||||
}
|
||||
PutFloat64Buf(next)
|
||||
}
|
||||
Reference in New Issue
Block a user