Files
tensor/internal/core/runtime_test.go
T
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

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)
}
}