From c234c3dd5b1e6e845eda25400e3fb53d2989e981 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Petr=20Balv=C3=ADn?= Date: Mon, 27 Jul 2026 21:52:16 +0200 Subject: [PATCH] feat(verify): add analyzeO1Range and fastStereoSums differential tests Assisted-by: Qwen 3.8 Max Preview --- cmd/gasm/main.go | 2 +- justfile | 2 +- verify/flac_test.go | 122 ++++++++++++++++++++++++++++++++++++++++++++ 3 files changed, 124 insertions(+), 2 deletions(-) diff --git a/cmd/gasm/main.go b/cmd/gasm/main.go index 440ba1b..ce86d6b 100644 --- a/cmd/gasm/main.go +++ b/cmd/gasm/main.go @@ -29,7 +29,7 @@ import ( // version is the release version, stamped at build time via // -ldflags "-X main.version=…" (defaulting to the current release). -var version = "0.21.0" +var version = "0.22.0" func main() { if len(os.Args) < 2 { diff --git a/justfile b/justfile index 66264e7..c75416c 100644 --- a/justfile +++ b/justfile @@ -3,7 +3,7 @@ # gasm-devkit — developer tooling for Go's Plan 9 assembler (GAsm). -version := "0.21.0" +version := "0.22.0" default: @just --list diff --git a/verify/flac_test.go b/verify/flac_test.go index 3a5c587..b4d3be1 100644 --- a/verify/flac_test.go +++ b/verify/flac_test.go @@ -75,6 +75,28 @@ func decorrelateInterleaveGo(left, right, out []int32) { } } +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 +} + // --- Differential tests --- func TestFLACDecodeMono16(t *testing.T) { @@ -206,3 +228,103 @@ func TestFLACDecorrelate(t *testing.T) { }) } } + +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)) +}