// 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] } } func analyzeO1RangeGo(swin []int32, dstP []uint32, hist *[32]uint16) (partSum uint64, overflow bool) { swin = swin[:len(dstP)+1] for j := 0; j+1 < len(swin); j++ { r := swin[j+1] - swin[j] if r == -2147483648 { // math.MinInt32 overflow = true } f := uint32(r<<1) ^ uint32(r>>31) dstP[j] = f partSum += uint64(f) bl := 0 for v := f; v > 0; v >>= 1 { bl++ } if bl > 31 { bl = 31 } hist[bl]++ } return } func analyzeO2RangeGo(swin []int32, dstP []uint32, hist *[32]uint16) (partSum uint64, overflow bool) { swin = swin[:len(dstP)+2] for j := 0; j+2 < len(swin); j++ { r := swin[j+2] - 2*swin[j+1] + swin[j] if r == -2147483648 { overflow = true } f := uint32(r<<1) ^ uint32(r>>31) dstP[j] = f partSum += uint64(f) bl := 0 for v := f; v > 0; v >>= 1 { bl++ } if bl > 31 { bl = 31 } hist[bl]++ } return } func analyzeResRangeGo(swin []int32, dstP []uint32, hist *[32]uint16) (partSum uint64, overflow bool) { for j := 0; j < len(swin); j++ { r := swin[j] if r == -2147483648 { overflow = true } f := uint32(r<<1) ^ uint32(r>>31) dstP[j] = f partSum += uint64(f) bl := 0 for v := f; v > 0; v >>= 1 { bl++ } if bl > 31 { bl = 31 } hist[bl]++ } return } func decodeMono24Go(src []byte, dst []int32) { for i := 0; i < len(dst); i++ { off := 3 * i u := uint32(src[off]) | uint32(src[off+1])<<8 | uint32(src[off+2])<<16 dst[i] = int32(u<<8) >> 8 } } func decodeStereo16Go(src []byte, left, right []int32) { for i := 0; i < len(left); i++ { left[i] = int32(int16(uint16(src[4*i]) | uint16(src[4*i+1])<<8)) right[i] = int32(int16(uint16(src[4*i+2]) | uint16(src[4*i+3])<<8)) } } // --- 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]) } } } }) } } func TestFLACAnalyzeO1Range(t *testing.T) { k := loadFLACKernel(t) rng := rand.New(rand.NewSource(33)) for iter := 0; iter < 300; iter++ { n := 1 + rng.Intn(128) // partition size swin := make([]int32, n+1) for i := range swin { swin[i] = int32(rng.Intn(1<<20) - 1<<19) } goDstP := make([]uint32, n) var goHist [32]uint16 goSum, goOvf := analyzeO1RangeGo(swin, goDstP, &goHist) jitDstP := make([]uint32, n) var jitHist [32]uint16 args := make([]byte, 72) // 65 rounded up PutPtr(args, 0, unsafe.Pointer(&swin[0])) PutUint64(args, 8, uint64(len(swin))) PutUint64(args, 16, uint64(cap(swin))) PutPtr(args, 24, unsafe.Pointer(&jitDstP[0])) PutUint64(args, 32, uint64(n)) PutUint64(args, 40, uint64(cap(jitDstP))) PutPtr(args, 48, unsafe.Pointer(&jitHist[0])) out, err := k.CallFunc("analyzeO1RangeAVX2", args) if err != nil { t.Fatalf("iter %d: %v", iter, err) } jitSum := GetUint64(out, 56) jitOvf := out[64] != 0 if jitSum != goSum { t.Fatalf("iter %d: partSum mismatch: JIT=%d Go=%d", iter, jitSum, goSum) } if jitOvf != goOvf { t.Fatalf("iter %d: overflow mismatch: JIT=%v Go=%v", iter, jitOvf, goOvf) } for i := range goDstP { if jitDstP[i] != goDstP[i] { 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) } } } func TestFLACFastStereoSums(t *testing.T) { k := loadFLACKernel(t) rng := rand.New(rand.NewSource(44)) for iter := 0; iter < 300; iter++ { n := 1 + rng.Intn(256) 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) } // Go reference: compute the four sums. var goSums [4]uint64 for i := 0; i < n; i++ { l := left[i] r := right[i] side := l - r mid := (l + r) >> 1 goSums[0] += foldAbs(l) + foldAbs(r) goSums[1] += foldAbs(l) + foldAbs(side) goSums[2] += foldAbs(side) + foldAbs(r) goSums[3] += foldAbs(mid) + foldAbs(side) } var jitSums [4]uint64 args := make([]byte, 56) PutPtr(args, 0, unsafe.Pointer(&left[0])) PutUint64(args, 8, uint64(n)) PutUint64(args, 16, uint64(cap(left))) PutPtr(args, 24, unsafe.Pointer(&right[0])) PutUint64(args, 32, uint64(n)) PutUint64(args, 40, uint64(cap(right))) PutPtr(args, 48, unsafe.Pointer(&jitSums[0])) _, err := k.CallFunc("fastStereoSumsAVX2", args) 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) } } } func foldAbs(v int32) uint64 { return uint64(uint32(v<<1) ^ uint32(v>>31)) } // runAnalyzeTest is the shared harness for the analyzeO*Range family. 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)) for iter := 0; iter < 200; iter++ { n := 1 + rng.Intn(128) swin := make([]int32, n+order) for i := range swin { swin[i] = int32(rng.Intn(1<<20) - 1<<19) } goDstP := make([]uint32, n) var goHist [32]uint16 goSum, goOvf := ref(swin, goDstP, &goHist) jitDstP := make([]uint32, n) var jitHist [32]uint16 args := make([]byte, 72) PutPtr(args, 0, unsafe.Pointer(&swin[0])) PutUint64(args, 8, uint64(len(swin))) PutUint64(args, 16, uint64(cap(swin))) PutPtr(args, 24, unsafe.Pointer(&jitDstP[0])) PutUint64(args, 32, uint64(n)) PutUint64(args, 40, uint64(cap(jitDstP))) PutPtr(args, 48, unsafe.Pointer(&jitHist[0])) out, err := k.CallFunc(name, args) if err != nil { t.Fatalf("iter %d: %v", iter, err) } jitSum := GetUint64(out, 56) jitOvf := out[64] != 0 if jitSum != goSum { t.Fatalf("iter %d: partSum: JIT=%d Go=%d", iter, jitSum, goSum) } if jitOvf != goOvf { t.Fatalf("iter %d: overflow: JIT=%v Go=%v", iter, jitOvf, goOvf) } for i := range goDstP { if jitDstP[i] != goDstP[i] { 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) } } } func TestFLACAnalyzeO2Range(t *testing.T) { k := loadFLACKernel(t) runAnalyzeTest(t, k, "analyzeO2RangeAVX2", 2, analyzeO2RangeGo) } func TestFLACAnalyzeResRange(t *testing.T) { k := loadFLACKernel(t) // analyzeResRange has order 0: swin IS the residual (no prediction). runAnalyzeTest(t, k, "analyzeResRangeAVX2", 0, analyzeResRangeGo) } func TestFLACAnalyzeO3Range(t *testing.T) { k := loadFLACKernel(t) ref := func(swin []int32, dstP []uint32, hist *[32]uint16) (uint64, bool) { swin = swin[:len(dstP)+3] var partSum uint64 var overflow bool for j := 0; j+3 < len(swin); j++ { r := swin[j+3] - 3*swin[j+2] + 3*swin[j+1] - swin[j] if r == -2147483648 { overflow = true } f := uint32(r<<1) ^ uint32(r>>31) dstP[j] = f partSum += uint64(f) bl := 0 for v := f; v > 0; v >>= 1 { bl++ } if bl > 31 { bl = 31 } hist[bl]++ } return partSum, overflow } runAnalyzeTest(t, k, "analyzeO3RangeAVX2", 3, ref) } func TestFLACAnalyzeO4Range(t *testing.T) { k := loadFLACKernel(t) ref := func(swin []int32, dstP []uint32, hist *[32]uint16) (uint64, bool) { swin = swin[:len(dstP)+4] var partSum uint64 var overflow bool for j := 0; j+4 < len(swin); j++ { r := swin[j+4] - 4*swin[j+3] + 6*swin[j+2] - 4*swin[j+1] + swin[j] if r == -2147483648 { overflow = true } f := uint32(r<<1) ^ uint32(r>>31) dstP[j] = f partSum += uint64(f) bl := 0 for v := f; v > 0; v >>= 1 { bl++ } if bl > 31 { bl = 31 } hist[bl]++ } return partSum, overflow } runAnalyzeTest(t, k, "analyzeO4RangeAVX2", 4, ref) } func TestFLACDecodeMono24(t *testing.T) { k := loadFLACKernel(t) rng := rand.New(rand.NewSource(55)) for iter := 0; iter < 500; iter++ { n := rng.Intn(256) src := make([]byte, 3*n) rng.Read(src) goDst := make([]int32, n) decodeMono24Go(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("decodeMono24AVX2", args) if err != nil { t.Fatalf("iter %d: %v", iter, err) } for i := range goDst { if jitDst[i] != goDst[i] { t.Fatalf("iter %d: dst[%d]: JIT=%d Go=%d", iter, i, jitDst[i], goDst[i]) } } } } func TestFLACDecodeStereo16(t *testing.T) { k := loadFLACKernel(t) rng := rand.New(rand.NewSource(66)) for iter := 0; iter < 500; iter++ { n := rng.Intn(256) src := make([]byte, 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) args := make([]byte, 72) 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(&jitLeft[0])) PutPtr(args, 48, unsafe.Pointer(&jitRight[0])) } PutUint64(args, 32, uint64(n)) PutUint64(args, 40, uint64(cap(jitLeft))) PutUint64(args, 56, uint64(n)) PutUint64(args, 64, uint64(cap(jitRight))) _, err := k.CallFunc("decodeStereo16AVX2", args) if err != nil { t.Fatalf("iter %d: %v", iter, err) } for i := 0; i < n; i++ { if jitLeft[i] != goLeft[i] || jitRight[i] != goRight[i] { t.Fatalf("iter %d: [%d] L: JIT=%d Go=%d; R: JIT=%d Go=%d", iter, i, jitLeft[i], goLeft[i], jitRight[i], goRight[i]) } } } }