90 lines
2.7 KiB
Go
90 lines
2.7 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|
// SPDX-License-Identifier: MIT
|
|
|
|
package core
|
|
|
|
import (
|
|
"runtime"
|
|
"testing"
|
|
)
|
|
|
|
func TestSetNumCPU(t *testing.T) {
|
|
// SetNumCPU(0) resets to runtime.NumCPU(): the invariant that must
|
|
// hold regardless of what previous tests set. Test against that,
|
|
// not against a captured "previous" value (tests run in any order).
|
|
ncpu := runtime.NumCPU()
|
|
if NumWorkers() < 1 {
|
|
t.Fatalf("NumWorkers: %d", NumWorkers())
|
|
}
|
|
got := SetNumCPU(4)
|
|
if got < 1 {
|
|
t.Errorf("SetNumCPU return: %d, want ≥ 1", got)
|
|
}
|
|
if NumWorkers() != 4 {
|
|
t.Errorf("NumWorkers after SetNumCPU(4): %d", NumWorkers())
|
|
}
|
|
// workersFor bounds against the explicit worker count, independent
|
|
// of the host's CPU count (the CI runner has one core).
|
|
if w := workersFor(2); w != 2 {
|
|
t.Errorf("workersFor(2) with 4 workers: %d, want 2", w)
|
|
}
|
|
if w := workersFor(1); w != 1 {
|
|
t.Errorf("workersFor(1): %d, want 1", w)
|
|
}
|
|
if w := workersFor(10); w != 4 {
|
|
t.Errorf("workersFor(10) with 4 workers: %d, want 4", w)
|
|
}
|
|
SetNumCPU(0) // resets to NumCPU
|
|
if NumWorkers() != ncpu {
|
|
t.Errorf("NumWorkers after reset: %d, want %d", NumWorkers(), ncpu)
|
|
}
|
|
// Universal invariants after the reset: at least one worker, never
|
|
// more than the item count.
|
|
if w := workersFor(0); w != 1 {
|
|
t.Errorf("workersFor(0): %d, want 1 (floor of one worker)", w)
|
|
}
|
|
if w := workersFor(1); w != 1 {
|
|
t.Errorf("workersFor(1) after reset: %d, want 1", w)
|
|
}
|
|
}
|
|
|
|
// TestParallelCoverage forces the parallel branch of `parallel` to run
|
|
// even on a single-core CI runner, where workersFor(n) would otherwise
|
|
// collapse to 1 and skip the goroutine-spawning path entirely, which
|
|
// silently drops the measured coverage below the gate. It also covers
|
|
// the parallel-only merge branch of the axis reductions (reduceAxis),
|
|
// whose private-scratch/merge code has no serial equivalent.
|
|
func TestParallelCoverage(t *testing.T) {
|
|
prev := SetNumCPU(2)
|
|
defer SetNumCPU(prev)
|
|
a, err := FromFloats(make([]float64, 1<<16), 1<<16)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
b, err := FromFloats(make([]float64, 1<<16), 1<<16)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
// With numWorkers=2 and 65536 items the parallel branch spawns
|
|
// goroutines even on a one-core host.
|
|
sum, err := Add(a, b)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if sum.Len() != a.Len() {
|
|
t.Fatalf("Add len: %d", sum.Len())
|
|
}
|
|
// Axis reductions take the parallel merge path with 2 workers,
|
|
// covering the private-scratch and merge code in reduceAxis.
|
|
mat, err := Reshape(a, 256, 256)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if _, err := SumAxis(mat, 1); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if _, err := MaxAxis(mat, 0); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
}
|