From a2301b52de10b4a0cff0c3150d9bb3cdcd3a7fce Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Petr=20Balv=C3=ADn?= Date: Tue, 6 Oct 2026 23:59:45 +0200 Subject: [PATCH] test(verify): run the amd64 extension encodings on the metal Assisted-by: GLM 5.3 Flash --- verify/amd64_ext_jit_test.go | 201 +++++++++++++++++++++++++++++++++++ 1 file changed, 201 insertions(+) create mode 100644 verify/amd64_ext_jit_test.go diff --git a/verify/amd64_ext_jit_test.go b/verify/amd64_ext_jit_test.go new file mode 100644 index 0000000..1a05dd0 --- /dev/null +++ b/verify/amd64_ext_jit_test.go @@ -0,0 +1,201 @@ +// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) +// SPDX-License-Identifier: BSD-3-Clause + +package verify + +import ( + "encoding/binary" + "fmt" + "math" + "os" + "runtime" + "strings" + "testing" + "unsafe" + + "sourcedock.dev/petrbalvin/gasm-sdk/arch" +) + +// The extended instructions have no toolchain oracle, so where the running +// CPU implements a family the layer's encodings are executed on the metal: +// the kernel below assembles through gasm (the baseline moves and the frame +// discipline) with the extension instruction laid byte for byte from the +// layer's own Encode output, and the result is checked against a portable Go +// reference of the manual's pseudo-code. On a CPU without the family the +// test skips: the golden vectors in the arch package are that path's proof. + +// requireCPUFlags skips unless the host lists every named CPUID flag. +func requireCPUFlags(t *testing.T, flags ...string) { + t.Helper() + if runtime.GOARCH != "amd64" || runtime.GOOS != "linux" { + t.Skipf("runs only on amd64 Linux hosts (this host is %s/%s)", runtime.GOOS, runtime.GOARCH) + } + data, err := os.ReadFile("/proc/cpuinfo") + if err != nil { + t.Skipf("cannot read the CPU flags: %v", err) + } + have := map[string]bool{} + for line := range strings.SplitSeq(string(data), "\n") { + if !strings.HasPrefix(line, "flags") { + continue + } + _, list, ok := strings.Cut(line, ":") + if !ok { + continue + } + for f := range strings.FieldsSeq(list) { + have[f] = true + } + } + for _, want := range flags { + if !have[want] { + t.Skipf("the CPU lacks %s", want) + } + } +} + +// amd64ExtEntry finds one encoding of one mnemonic at the 512-bit length. +func amd64ExtEntry(t *testing.T, mnem string, ops ...arch.ExtOperand) []byte { + t.Helper() + for _, in := range arch.Extensions(arch.AMD64) { + if in.Name == mnem && in.Bytes[3]>>5&3 == 2 { + b, err := in.Encode(ops) + if err != nil { + t.Fatalf("%s: encode: %v", mnem, err) + } + return b + } + } + t.Fatalf("the layer registers no 512-bit %s", mnem) + return nil +} + +// extByteLines renders an encoding as BYTE lines the assembler lays verbatim, +// the Plan 9 way of naming machine bytes the instruction table lacks. +func extByteLines(b []byte) string { + var sb strings.Builder + for _, x := range b { + fmt.Fprintf(&sb, "\tBYTE $0x%02x\n", x) + } + return sb.String() +} + +// bf16Round rounds a float32 bit pattern to BF16, nearest even: the manual's +// VCVTNEPS2BF16 carries the NE of no exception, not of truncation, so the +// low sixteen mantissa bits round and carry into the exponent. +func bf16Round(bits uint32) uint16 { + bias := uint32(0x7fff) + bits>>16&1 + return uint16((bits + bias) >> 16) +} + +// TestJITAmd64ExtBF16 converts sixteen float32 values to BF16 with the +// layer's VCVTNEPS2BF16 encoding and checks the result against the manual's +// rounding: nearest even, no FP exception. +func TestJITAmd64ExtBF16(t *testing.T) { + requireCPUFlags(t, "avx512f", "avx512_bf16") + + ext := amd64ExtEntry(t, "VCVTNEPS2BF16", arch.ExtZmm(0), arch.ExtYmm(1)) + src := "#include \"textflag.h\"\n" + ` +// func cvtbf16(p, q *byte) +TEXT ·cvtbf16(SB), NOSPLIT, $0-16 + MOVQ p+0(FP), SI + MOVQ q+8(FP), DI + VMOVUPS (SI), Z0 +` + extByteLines(ext) + ` VMOVUPS Y1, (DI) + VZEROUPPER + RET +` + k, err := LoadSource("amd64_ext_bf16.s", src) + if err != nil { + t.Fatalf("LoadSource: %v", err) + } + t.Cleanup(k.Close) + + in := []float32{1.0, -2.5, 0.0, math.Pi, 1e10, -0.5, 65504, 1e-10, + -1.0, 2.5, 1024.0, 0.25, 1e20, -3.0, 0.5, 9.75} + out := make([]byte, 32) + args := make([]byte, 16) + PutPtr(args, 0, unsafe.Pointer(&in[0])) + PutPtr(args, 8, unsafe.Pointer(&out[0])) + if _, err := k.CallFunc("cvtbf16", args); err != nil { + t.Fatalf("CallFunc: %v", err) + } + for i, f := range in { + want := bf16Round(math.Float32bits(f)) + if got := uint16(out[2*i]) | uint16(out[2*i+1])<<8; got != want { + t.Errorf("bf16(%v) = %#04x, want %#04x", f, got, want) + } + } +} + +// TestJITAmd64ExtVP2INTERSECT intersects two dword vectors with the layer's +// VP2INTERSECTD encoding and checks both halves against the manual: the +// destination is an even/odd mask register pair, the even register marking +// the first source's elements found in the second, the odd one the second +// source's elements found in the first. +func TestJITAmd64ExtVP2INTERSECT(t *testing.T) { + requireCPUFlags(t, "avx512f", "avx512_vp2intersect") + + ext := amd64ExtEntry(t, "VP2INTERSECTD", arch.ExtZmm(0), arch.ExtZmm(1), arch.ExtMask(0)) + src := "#include \"textflag.h\"\n" + ` +// func isect(p, q, r *byte) +TEXT ·isect(SB), NOSPLIT, $0-24 + MOVQ p+0(FP), SI + MOVQ q+8(FP), DI + MOVQ r+16(FP), DX + VMOVUPS (SI), Z0 + VMOVUPS (DI), Z1 +` + extByteLines(ext) + ` KMOVD K0, AX + KMOVD K1, CX + MOVL AX, (DX) + MOVL CX, 4(DX) + VZEROUPPER + RET +` + k, err := LoadSource("amd64_ext_vp2intersect.s", src) + if err != nil { + t.Fatalf("LoadSource: %v", err) + } + t.Cleanup(k.Close) + + a := []uint32{10, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24} + b := []uint32{10, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30, 31, 32, 33, 34} + var lo, hi uint16 + for i := range 16 { + for j := range 16 { + if a[i] == b[j] { + lo |= 1 << i + break + } + } + } + for j := range 16 { + for i := range 16 { + if b[j] == a[i] { + hi |= 1 << j + break + } + } + } + want := uint64(lo) | uint64(hi)<<32 + + bufA := make([]byte, 64) + bufB := make([]byte, 64) + for i, v := range a { + binary.LittleEndian.PutUint32(bufA[4*i:], v) + } + for i, v := range b { + binary.LittleEndian.PutUint32(bufB[4*i:], v) + } + out := make([]byte, 8) + args := make([]byte, 24) + PutPtr(args, 0, unsafe.Pointer(&bufA[0])) + PutPtr(args, 8, unsafe.Pointer(&bufB[0])) + PutPtr(args, 16, unsafe.Pointer(&out[0])) + if _, err := k.CallFunc("isect", args); err != nil { + t.Fatalf("CallFunc: %v", err) + } + if got := binary.LittleEndian.Uint64(out); got != want { + t.Errorf("the intersection masks are %#018x, want %#018x", got, want) + } +}