65 lines
1.9 KiB
Go
65 lines
1.9 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|
// SPDX-License-Identifier: MIT
|
|
|
|
package integrate
|
|
|
|
import (
|
|
"math"
|
|
"testing"
|
|
|
|
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
|
"sourcedock.dev/petrbalvin/tensor/internal/engine"
|
|
)
|
|
|
|
// The 2-D step sweeps index their per-line tridiagonal scratch by
|
|
// chunk start divided by the spawn floor, which is only exercised when
|
|
// the engine actually spawns: on a many-core runner the fixture grids
|
|
// run inline and a scratch-indexing break passes silently. This pin
|
|
// forces both regimes and compares every output bit.
|
|
|
|
func pinHeat2D(workers int) ([]float64, error) {
|
|
prev := engine.SetNumWorkers(workers)
|
|
defer engine.SetNumWorkers(prev)
|
|
const rows, cols = 20, 20
|
|
u0 := make([]float64, rows*cols)
|
|
for r := range rows {
|
|
for c := range cols {
|
|
u0[r*cols+c] = math.Sin(float64(c+1)/float64(cols+1)*math.Pi) *
|
|
math.Sin(float64(r+1)/float64(rows+1)*math.Pi)
|
|
}
|
|
}
|
|
state, err := core.FromFloats(u0, rows, cols)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
out, err := IntegrateHeat2D(state, 1, 1.0/21, 1.0/21, 0.02, 0.002, 2, 0, 0, 0, 0)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return append([]float64{}, out.RawFloats()...), nil
|
|
}
|
|
|
|
func TestHeat2DScratchIndexingPinnedAcrossWorkers(t *testing.T) {
|
|
serial, err := pinHeat2D(1)
|
|
if err != nil {
|
|
t.Fatalf("serial: %v", err)
|
|
}
|
|
// Two workers spawn on an 18-line sweep (chunk 9 at the floor of
|
|
// 8); four and thirty-two collapse back to inline runs.
|
|
for _, w := range []int{2, 4, 32} {
|
|
got, err := pinHeat2D(w)
|
|
if err != nil {
|
|
t.Fatalf("workers=%d: %v", w, err)
|
|
}
|
|
if len(got) != len(serial) {
|
|
t.Fatalf("workers=%d: length %d, want %d", w, len(got), len(serial))
|
|
}
|
|
for i := range serial {
|
|
if math.Float64bits(got[i]) != math.Float64bits(serial[i]) {
|
|
t.Fatalf("workers=%d element %d: %#x, want %#x",
|
|
w, i, math.Float64bits(got[i]), math.Float64bits(serial[i]))
|
|
}
|
|
}
|
|
}
|
|
}
|