fix(verify): fix flaky JIT tests with global buffers and KeepAlive
This commit is contained in:
+134
-27
@@ -7,6 +7,7 @@ import (
|
||||
"bytes"
|
||||
"math/rand"
|
||||
"os"
|
||||
"runtime"
|
||||
"testing"
|
||||
"unsafe"
|
||||
)
|
||||
@@ -157,19 +158,51 @@ func decodeStereo16Go(src []byte, left, right []int32) {
|
||||
|
||||
// --- Differential tests ---
|
||||
|
||||
// Buffers whose addresses are passed to JIT code via unsafe.Pointer are
|
||||
// package-level globals: their addresses are stable and the GC never moves
|
||||
// or collects them, unlike per-iteration make() buffers (which caused flaky
|
||||
// stale reads, especially under -race). Sizes cover the maximum n from the
|
||||
// rng.Intn(N) call in each test, rounded up.
|
||||
var (
|
||||
mono16Src [512]byte // TestFLACDecodeMono16: max n = 255 → 2*n = 510
|
||||
mono16JitDst [256]int32 // max n = 255
|
||||
|
||||
pack16Src [256]int32 // TestFLACPack16: max n = 255
|
||||
pack16JitDst [512]byte // max 2*n = 510
|
||||
|
||||
decorLeft [128]int32 // TestFLACDecorrelate: max n = 127
|
||||
decorRight [128]int32 // max n = 127
|
||||
decorJitOut [256]int32 // max 2*n = 254
|
||||
|
||||
mono24Src [768]byte // TestFLACDecodeMono24: max n = 255 → 3*n = 765
|
||||
mono24JitDst [256]int32 // max n = 255
|
||||
|
||||
stereo16Src [1024]byte // TestFLACDecodeStereo16: max n = 255 → 4*n = 1020
|
||||
stereo16JitLeft [256]int32 // max n = 255
|
||||
stereo16JitRight [256]int32 // max n = 255
|
||||
|
||||
fastStereoLeft [256]int32 // TestFLACFastStereoSums: max n = 255
|
||||
fastStereoRight [256]int32 // max n = 255
|
||||
)
|
||||
|
||||
func TestFLACDecodeMono16(t *testing.T) {
|
||||
k := loadFLACKernel(t)
|
||||
rng := rand.New(rand.NewSource(7))
|
||||
|
||||
src := mono16Src[:0]
|
||||
jitDst := mono16JitDst[:0]
|
||||
for iter := 0; iter < 500; iter++ {
|
||||
n := rng.Intn(256)
|
||||
src := make([]byte, 2*n)
|
||||
src = src[:2*n]
|
||||
rng.Read(src)
|
||||
|
||||
goDst := make([]int32, n)
|
||||
decodeMono16Go(src, goDst)
|
||||
|
||||
jitDst := make([]int32, n)
|
||||
jitDst = jitDst[:n]
|
||||
for i := range jitDst {
|
||||
jitDst[i] = 0
|
||||
}
|
||||
args := make([]byte, 48)
|
||||
if len(src) > 0 {
|
||||
PutPtr(args, 0, unsafe.Pointer(&src[0]))
|
||||
@@ -183,6 +216,8 @@ func TestFLACDecodeMono16(t *testing.T) {
|
||||
PutUint64(args, 40, uint64(cap(jitDst)))
|
||||
|
||||
_, err := k.CallFunc("decodeMono16AVX2", args)
|
||||
runtime.KeepAlive(src)
|
||||
runtime.KeepAlive(jitDst)
|
||||
if err != nil {
|
||||
t.Fatalf("iter %d: %v", iter, err)
|
||||
}
|
||||
@@ -198,9 +233,11 @@ func TestFLACPack16(t *testing.T) {
|
||||
k := loadFLACKernel(t)
|
||||
rng := rand.New(rand.NewSource(13))
|
||||
|
||||
src := pack16Src[:0]
|
||||
jitDst := pack16JitDst[:0]
|
||||
for iter := 0; iter < 500; iter++ {
|
||||
n := rng.Intn(256)
|
||||
src := make([]int32, n)
|
||||
src = src[:n]
|
||||
for i := range src {
|
||||
src[i] = int32(rng.Intn(65536) - 32768)
|
||||
}
|
||||
@@ -208,7 +245,10 @@ func TestFLACPack16(t *testing.T) {
|
||||
goDst := make([]byte, 2*n)
|
||||
pack16Go(goDst, src)
|
||||
|
||||
jitDst := make([]byte, 2*n)
|
||||
jitDst = jitDst[:2*n]
|
||||
for i := range jitDst {
|
||||
jitDst[i] = 0
|
||||
}
|
||||
args := make([]byte, 48)
|
||||
if len(jitDst) > 0 {
|
||||
PutPtr(args, 0, unsafe.Pointer(&jitDst[0]))
|
||||
@@ -222,6 +262,8 @@ func TestFLACPack16(t *testing.T) {
|
||||
PutUint64(args, 40, uint64(cap(src)))
|
||||
|
||||
_, err := k.CallFunc("pack16AVX2", args)
|
||||
runtime.KeepAlive(jitDst)
|
||||
runtime.KeepAlive(src)
|
||||
if err != nil {
|
||||
t.Fatalf("iter %d: %v", iter, err)
|
||||
}
|
||||
@@ -247,10 +289,13 @@ func TestFLACDecorrelate(t *testing.T) {
|
||||
|
||||
for _, kk := range kernels {
|
||||
t.Run(kk.name, func(t *testing.T) {
|
||||
left := decorLeft[:0]
|
||||
right := decorRight[:0]
|
||||
jitOut := decorJitOut[:0]
|
||||
for iter := 0; iter < 200; iter++ {
|
||||
n := rng.Intn(128)
|
||||
left := make([]int32, n)
|
||||
right := make([]int32, n)
|
||||
left = left[:n]
|
||||
right = right[:n]
|
||||
for i := range left {
|
||||
left[i] = int32(rng.Intn(1<<24) - 1<<23)
|
||||
right[i] = int32(rng.Intn(1<<24) - 1<<23)
|
||||
@@ -259,7 +304,10 @@ func TestFLACDecorrelate(t *testing.T) {
|
||||
goOut := make([]int32, 2*n)
|
||||
kk.ref(left, right, goOut)
|
||||
|
||||
jitOut := make([]int32, 2*n)
|
||||
jitOut = jitOut[:2*n]
|
||||
for i := range jitOut {
|
||||
jitOut[i] = 0
|
||||
}
|
||||
args := make([]byte, 72)
|
||||
if n > 0 {
|
||||
PutPtr(args, 0, unsafe.Pointer(&left[0]))
|
||||
@@ -274,6 +322,9 @@ func TestFLACDecorrelate(t *testing.T) {
|
||||
PutUint64(args, 64, uint64(cap(jitOut)))
|
||||
|
||||
_, err := k.CallFunc(kk.name, args)
|
||||
runtime.KeepAlive(left)
|
||||
runtime.KeepAlive(right)
|
||||
runtime.KeepAlive(jitOut)
|
||||
if err != nil {
|
||||
t.Fatalf("iter %d: %v", iter, err)
|
||||
}
|
||||
@@ -291,9 +342,11 @@ func TestFLACAnalyzeO1Range(t *testing.T) {
|
||||
k := loadFLACKernel(t)
|
||||
rng := rand.New(rand.NewSource(33))
|
||||
|
||||
swin := analyzeSwin[:0]
|
||||
jitDstP := analyzeJitDstP[:0]
|
||||
for iter := 0; iter < 300; iter++ {
|
||||
n := 1 + rng.Intn(128) // partition size
|
||||
swin := make([]int32, n+1)
|
||||
swin = swin[:n+1]
|
||||
for i := range swin {
|
||||
swin[i] = int32(rng.Intn(1<<20) - 1<<19)
|
||||
}
|
||||
@@ -302,8 +355,12 @@ func TestFLACAnalyzeO1Range(t *testing.T) {
|
||||
var goHist [32]uint16
|
||||
goSum, goOvf := analyzeO1RangeGo(swin, goDstP, &goHist)
|
||||
|
||||
jitDstP := make([]uint32, n)
|
||||
var jitHist [32]uint16
|
||||
jitDstP = jitDstP[:n]
|
||||
for i := range jitDstP {
|
||||
jitDstP[i] = 0
|
||||
}
|
||||
jitHist := &analyzeJitHist
|
||||
*jitHist = [32]uint16{} // reset for this iteration
|
||||
args := make([]byte, 72) // 65 rounded up
|
||||
PutPtr(args, 0, unsafe.Pointer(&swin[0]))
|
||||
PutUint64(args, 8, uint64(len(swin)))
|
||||
@@ -314,6 +371,9 @@ func TestFLACAnalyzeO1Range(t *testing.T) {
|
||||
PutPtr(args, 48, unsafe.Pointer(&jitHist[0]))
|
||||
|
||||
out, err := k.CallFunc("analyzeO1RangeAVX2", args)
|
||||
runtime.KeepAlive(swin)
|
||||
runtime.KeepAlive(jitDstP)
|
||||
runtime.KeepAlive(jitHist)
|
||||
if err != nil {
|
||||
t.Fatalf("iter %d: %v", iter, err)
|
||||
}
|
||||
@@ -331,8 +391,8 @@ func TestFLACAnalyzeO1Range(t *testing.T) {
|
||||
t.Fatalf("iter %d: dstP[%d] mismatch: JIT=%d Go=%d", iter, i, jitDstP[i], goDstP[i])
|
||||
}
|
||||
}
|
||||
if jitHist != goHist {
|
||||
t.Fatalf("iter %d: hist mismatch: JIT=%v Go=%v", iter, jitHist, goHist)
|
||||
if *jitHist != goHist {
|
||||
t.Fatalf("iter %d: hist mismatch: JIT=%v Go=%v", iter, *jitHist, goHist)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -341,10 +401,12 @@ func TestFLACFastStereoSums(t *testing.T) {
|
||||
k := loadFLACKernel(t)
|
||||
rng := rand.New(rand.NewSource(44))
|
||||
|
||||
left := fastStereoLeft[:0]
|
||||
right := fastStereoRight[:0]
|
||||
for iter := 0; iter < 300; iter++ {
|
||||
n := 1 + rng.Intn(256)
|
||||
left := make([]int32, n)
|
||||
right := make([]int32, n)
|
||||
left = left[:n]
|
||||
right = right[:n]
|
||||
for i := range left {
|
||||
left[i] = int32(rng.Intn(1<<24) - 1<<23)
|
||||
right[i] = int32(rng.Intn(1<<24) - 1<<23)
|
||||
@@ -363,7 +425,8 @@ func TestFLACFastStereoSums(t *testing.T) {
|
||||
goSums[3] += foldAbs(mid) + foldAbs(side)
|
||||
}
|
||||
|
||||
var jitSums [4]uint64
|
||||
jitSums := &analyzeJitSums
|
||||
*jitSums = [4]uint64{} // reset for this iteration
|
||||
args := make([]byte, 56)
|
||||
PutPtr(args, 0, unsafe.Pointer(&left[0]))
|
||||
PutUint64(args, 8, uint64(n))
|
||||
@@ -374,11 +437,14 @@ func TestFLACFastStereoSums(t *testing.T) {
|
||||
PutPtr(args, 48, unsafe.Pointer(&jitSums[0]))
|
||||
|
||||
_, err := k.CallFunc("fastStereoSumsAVX2", args)
|
||||
runtime.KeepAlive(left)
|
||||
runtime.KeepAlive(right)
|
||||
runtime.KeepAlive(jitSums)
|
||||
if err != nil {
|
||||
t.Fatalf("iter %d: %v", iter, err)
|
||||
}
|
||||
if jitSums != goSums {
|
||||
t.Fatalf("iter %d: sums mismatch:\n JIT=%v\n Go =%v", iter, jitSums, goSums)
|
||||
if *jitSums != goSums {
|
||||
t.Fatalf("iter %d: sums mismatch:\n JIT=%v\n Go =%v", iter, *jitSums, goSums)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -388,12 +454,27 @@ func foldAbs(v int32) uint64 {
|
||||
}
|
||||
|
||||
// runAnalyzeTest is the shared harness for the analyzeO*Range family.
|
||||
// analyzeJitHist is a shared histogram buffer for the analyze tests.
|
||||
// It is package-level so its address is stable across all calls.
|
||||
var analyzeJitHist [32]uint16
|
||||
|
||||
// analyzeJitSums is a shared sums buffer for the fastStereoSums test.
|
||||
var analyzeJitSums [4]uint64
|
||||
|
||||
// analyzeSwin and analyzeJitDstP are shared input/output buffers for the
|
||||
// analyze tests. Package-level so their addresses are stable.
|
||||
var analyzeSwin [136]int32 // max n(128) + max order(4) = 132, rounded up
|
||||
var analyzeJitDstP [128]uint32
|
||||
|
||||
func runAnalyzeTest(t *testing.T, k *Kernel, name string, order int, ref func([]int32, []uint32, *[32]uint16) (uint64, bool)) {
|
||||
t.Helper()
|
||||
rng := rand.New(rand.NewSource(int64(order)*100 + 7))
|
||||
jitHist := &analyzeJitHist
|
||||
swin := analyzeSwin[:0]
|
||||
jitDstP := analyzeJitDstP[:0]
|
||||
for iter := 0; iter < 200; iter++ {
|
||||
n := 1 + rng.Intn(128)
|
||||
swin := make([]int32, n+order)
|
||||
swin = swin[:n+order]
|
||||
for i := range swin {
|
||||
swin[i] = int32(rng.Intn(1<<20) - 1<<19)
|
||||
}
|
||||
@@ -402,8 +483,11 @@ func runAnalyzeTest(t *testing.T, k *Kernel, name string, order int, ref func([]
|
||||
var goHist [32]uint16
|
||||
goSum, goOvf := ref(swin, goDstP, &goHist)
|
||||
|
||||
jitDstP := make([]uint32, n)
|
||||
var jitHist [32]uint16
|
||||
jitDstP = jitDstP[:n]
|
||||
for i := range jitDstP {
|
||||
jitDstP[i] = 0
|
||||
}
|
||||
*jitHist = [32]uint16{} // reset for this iteration
|
||||
args := make([]byte, 72)
|
||||
PutPtr(args, 0, unsafe.Pointer(&swin[0]))
|
||||
PutUint64(args, 8, uint64(len(swin)))
|
||||
@@ -413,7 +497,13 @@ func runAnalyzeTest(t *testing.T, k *Kernel, name string, order int, ref func([]
|
||||
PutUint64(args, 40, uint64(cap(jitDstP)))
|
||||
PutPtr(args, 48, unsafe.Pointer(&jitHist[0]))
|
||||
|
||||
// Memory barrier: ensure all writes above are visible to the JIT code.
|
||||
runtime.KeepAlive(args)
|
||||
|
||||
out, err := k.CallFunc(name, args)
|
||||
runtime.KeepAlive(swin)
|
||||
runtime.KeepAlive(jitDstP)
|
||||
runtime.KeepAlive(jitHist)
|
||||
if err != nil {
|
||||
t.Fatalf("iter %d: %v", iter, err)
|
||||
}
|
||||
@@ -431,8 +521,8 @@ func runAnalyzeTest(t *testing.T, k *Kernel, name string, order int, ref func([]
|
||||
t.Fatalf("iter %d: dstP[%d]: JIT=%d Go=%d", iter, i, jitDstP[i], goDstP[i])
|
||||
}
|
||||
}
|
||||
if jitHist != goHist {
|
||||
t.Fatalf("iter %d: hist mismatch", iter)
|
||||
if *jitHist != goHist {
|
||||
t.Fatalf("iter %d: hist mismatch\n JIT=%v\n Go =%v\n n=%d", iter, *jitHist, goHist, n)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -508,15 +598,20 @@ func TestFLACDecodeMono24(t *testing.T) {
|
||||
k := loadFLACKernel(t)
|
||||
rng := rand.New(rand.NewSource(55))
|
||||
|
||||
src := mono24Src[:0]
|
||||
jitDst := mono24JitDst[:0]
|
||||
for iter := 0; iter < 500; iter++ {
|
||||
n := rng.Intn(256)
|
||||
src := make([]byte, 3*n)
|
||||
src = src[:3*n]
|
||||
rng.Read(src)
|
||||
|
||||
goDst := make([]int32, n)
|
||||
decodeMono24Go(src, goDst)
|
||||
|
||||
jitDst := make([]int32, n)
|
||||
jitDst = jitDst[:n]
|
||||
for i := range jitDst {
|
||||
jitDst[i] = 0
|
||||
}
|
||||
args := make([]byte, 48)
|
||||
if len(src) > 0 {
|
||||
PutPtr(args, 0, unsafe.Pointer(&src[0]))
|
||||
@@ -530,6 +625,8 @@ func TestFLACDecodeMono24(t *testing.T) {
|
||||
PutUint64(args, 40, uint64(cap(jitDst)))
|
||||
|
||||
_, err := k.CallFunc("decodeMono24AVX2", args)
|
||||
runtime.KeepAlive(src)
|
||||
runtime.KeepAlive(jitDst)
|
||||
if err != nil {
|
||||
t.Fatalf("iter %d: %v", iter, err)
|
||||
}
|
||||
@@ -545,17 +642,24 @@ func TestFLACDecodeStereo16(t *testing.T) {
|
||||
k := loadFLACKernel(t)
|
||||
rng := rand.New(rand.NewSource(66))
|
||||
|
||||
src := stereo16Src[:0]
|
||||
jitLeft := stereo16JitLeft[:0]
|
||||
jitRight := stereo16JitRight[:0]
|
||||
for iter := 0; iter < 500; iter++ {
|
||||
n := rng.Intn(256)
|
||||
src := make([]byte, 4*n) // [L0,R0,L1,R1,...]
|
||||
src = src[:4*n] // [L0,R0,L1,R1,...]
|
||||
rng.Read(src)
|
||||
|
||||
goLeft := make([]int32, n)
|
||||
goRight := make([]int32, n)
|
||||
decodeStereo16Go(src, goLeft, goRight)
|
||||
|
||||
jitLeft := make([]int32, n)
|
||||
jitRight := make([]int32, n)
|
||||
jitLeft = jitLeft[:n]
|
||||
jitRight = jitRight[:n]
|
||||
for i := range jitLeft {
|
||||
jitLeft[i] = 0
|
||||
jitRight[i] = 0
|
||||
}
|
||||
args := make([]byte, 72)
|
||||
if len(src) > 0 {
|
||||
PutPtr(args, 0, unsafe.Pointer(&src[0]))
|
||||
@@ -572,6 +676,9 @@ func TestFLACDecodeStereo16(t *testing.T) {
|
||||
PutUint64(args, 64, uint64(cap(jitRight)))
|
||||
|
||||
_, err := k.CallFunc("decodeStereo16AVX2", args)
|
||||
runtime.KeepAlive(src)
|
||||
runtime.KeepAlive(jitLeft)
|
||||
runtime.KeepAlive(jitRight)
|
||||
if err != nil {
|
||||
t.Fatalf("iter %d: %v", iter, err)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user