// Copyright (c) 2026 Petr BalvĂ­n (https://petrbalvin.org) // SPDX-License-Identifier: BSD-3-Clause package verify import ( "bytes" "math/rand" "os" "testing" "unsafe" ) const flacKernelPath = "../../go-libraries/go-flac/avx2_amd64.s" func loadFLACKernel(t *testing.T) *Kernel { t.Helper() if _, err := os.Stat(flacKernelPath); err != nil { t.Skipf("sibling kernel not available: %v", err) } k, err := Load(flacKernelPath) if err != nil { t.Fatalf("Load(%s): %v", flacKernelPath, err) } t.Cleanup(k.Close) return k } // --- Portable Go references (from go-flac/simd.go) --- func decodeMono16Go(src []byte, dst []int32) { for i := 0; i < len(dst); i++ { dst[i] = int32(int16(uint16(src[2*i]) | uint16(src[2*i+1])<<8)) } } func pack16Go(dst []byte, src []int32) { for i, v := range src { dst[2*i] = byte(v) dst[2*i+1] = byte(v >> 8) } } func decorrelateLeftSideGo(left, right, out []int32) { for i := range left { l := left[i] out[2*i] = l out[2*i+1] = l - right[i] } } func decorrelateSideRightGo(left, right, out []int32) { for i := range left { side := left[i] rch := right[i] out[2*i] = rch + side out[2*i+1] = rch } } func decorrelateMidSideGo(left, right, out []int32) { for i := range left { mid := left[i] side := right[i] mid2 := mid<<1 | (side & 1) out[2*i] = (mid2 + side) >> 1 out[2*i+1] = (mid2 - side) >> 1 } } func decorrelateInterleaveGo(left, right, out []int32) { for i := range left { out[2*i] = left[i] out[2*i+1] = right[i] } } // --- Differential tests --- func TestFLACDecodeMono16(t *testing.T) { k := loadFLACKernel(t) rng := rand.New(rand.NewSource(7)) for iter := 0; iter < 500; iter++ { n := rng.Intn(256) src := make([]byte, 2*n) rng.Read(src) goDst := make([]int32, n) decodeMono16Go(src, goDst) jitDst := make([]int32, n) args := make([]byte, 48) if len(src) > 0 { PutPtr(args, 0, unsafe.Pointer(&src[0])) } PutUint64(args, 8, uint64(len(src))) PutUint64(args, 16, uint64(cap(src))) if n > 0 { PutPtr(args, 24, unsafe.Pointer(&jitDst[0])) } PutUint64(args, 32, uint64(n)) PutUint64(args, 40, uint64(cap(jitDst))) _, err := k.CallFunc("decodeMono16AVX2", args) if err != nil { t.Fatalf("iter %d: %v", iter, err) } for i := range goDst { if jitDst[i] != goDst[i] { t.Fatalf("iter %d: mismatch at [%d]: JIT=%d Go=%d", iter, i, jitDst[i], goDst[i]) } } } } func TestFLACPack16(t *testing.T) { k := loadFLACKernel(t) rng := rand.New(rand.NewSource(13)) for iter := 0; iter < 500; iter++ { n := rng.Intn(256) src := make([]int32, n) for i := range src { src[i] = int32(rng.Intn(65536) - 32768) } goDst := make([]byte, 2*n) pack16Go(goDst, src) jitDst := make([]byte, 2*n) args := make([]byte, 48) if len(jitDst) > 0 { PutPtr(args, 0, unsafe.Pointer(&jitDst[0])) } PutUint64(args, 8, uint64(len(jitDst))) PutUint64(args, 16, uint64(cap(jitDst))) if n > 0 { PutPtr(args, 24, unsafe.Pointer(&src[0])) } PutUint64(args, 32, uint64(n)) PutUint64(args, 40, uint64(cap(src))) _, err := k.CallFunc("pack16AVX2", args) if err != nil { t.Fatalf("iter %d: %v", iter, err) } if !bytes.Equal(jitDst, goDst) { t.Fatalf("iter %d: output mismatch (n=%d)", iter, n) } } } func TestFLACDecorrelate(t *testing.T) { k := loadFLACKernel(t) rng := rand.New(rand.NewSource(21)) kernels := []struct { name string ref func(left, right, out []int32) }{ {"decorrelateLeftSideAVX2", decorrelateLeftSideGo}, {"decorrelateSideRightAVX2", decorrelateSideRightGo}, {"decorrelateMidSideAVX2", decorrelateMidSideGo}, {"decorrelateInterleaveAVX2", decorrelateInterleaveGo}, } for _, kk := range kernels { t.Run(kk.name, func(t *testing.T) { for iter := 0; iter < 200; iter++ { n := rng.Intn(128) left := make([]int32, n) right := make([]int32, n) for i := range left { left[i] = int32(rng.Intn(1<<24) - 1<<23) right[i] = int32(rng.Intn(1<<24) - 1<<23) } goOut := make([]int32, 2*n) kk.ref(left, right, goOut) jitOut := make([]int32, 2*n) args := make([]byte, 72) if n > 0 { PutPtr(args, 0, unsafe.Pointer(&left[0])) PutPtr(args, 24, unsafe.Pointer(&right[0])) PutPtr(args, 48, unsafe.Pointer(&jitOut[0])) } PutUint64(args, 8, uint64(n)) PutUint64(args, 16, uint64(cap(left))) PutUint64(args, 32, uint64(n)) PutUint64(args, 40, uint64(cap(right))) PutUint64(args, 56, uint64(2*n)) PutUint64(args, 64, uint64(cap(jitOut))) _, err := k.CallFunc(kk.name, args) if err != nil { t.Fatalf("iter %d: %v", iter, err) } for i := range goOut { if jitOut[i] != goOut[i] { t.Fatalf("iter %d: mismatch at [%d]: JIT=%d Go=%d", iter, i, jitOut[i], goOut[i]) } } } }) } }