2026-10-06 23:59:45 +02:00
|
|
|
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (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)
|
|
|
|
|
}
|
|
|
|
|
}
|
2026-10-07 00:04:33 +02:00
|
|
|
|
|
|
|
|
// TestJITAmd64ExtVDPBF16PS accumulates the BF16 dot product with the layer's
|
|
|
|
|
// VDPBF16PS encoding. The lanes hold dyadic BF16 values whose products and
|
|
|
|
|
// sums are exact in float32, so the reference is independent of the rounding
|
|
|
|
|
// order and the check pins the lane layout: each dword lane of the
|
|
|
|
|
// accumulator takes the high halves' product plus the low halves' product.
|
|
|
|
|
func TestJITAmd64ExtVDPBF16PS(t *testing.T) {
|
|
|
|
|
requireCPUFlags(t, "avx512f", "avx512_bf16")
|
|
|
|
|
|
|
|
|
|
ext := amd64ExtEntry(t, "VDPBF16PS", arch.ExtZmm(0), arch.ExtZmm(1), arch.ExtZmm(2))
|
|
|
|
|
src := "#include \"textflag.h\"\n" + `
|
|
|
|
|
// func dp(a, b, c *byte)
|
|
|
|
|
TEXT ·dp(SB), NOSPLIT, $0-24
|
|
|
|
|
MOVQ a+0(FP), SI
|
|
|
|
|
MOVQ b+8(FP), DI
|
|
|
|
|
MOVQ c+16(FP), DX
|
|
|
|
|
VMOVUPS (SI), Z0
|
|
|
|
|
VMOVUPS (DI), Z1
|
|
|
|
|
VMOVUPS (DX), Z2
|
|
|
|
|
` + extByteLines(ext) + ` VMOVUPS Z2, (DX)
|
|
|
|
|
VZEROUPPER
|
|
|
|
|
RET
|
|
|
|
|
`
|
|
|
|
|
k, err := LoadSource("amd64_ext_vdpbf16ps.s", src)
|
|
|
|
|
if err != nil {
|
|
|
|
|
t.Fatalf("LoadSource: %v", err)
|
|
|
|
|
}
|
|
|
|
|
t.Cleanup(k.Close)
|
|
|
|
|
|
|
|
|
|
// Each dword lane pairs two BF16 halves; the table gives the halves as
|
|
|
|
|
// float32 values with an exact BF16 representation.
|
|
|
|
|
pairs := [][4]float32{
|
|
|
|
|
{2.0, 3.0, 1.5, 2.0}, // a high, b high, a low, b low
|
|
|
|
|
{0.5, 4.0, 1.0, 2.5},
|
|
|
|
|
}
|
|
|
|
|
aBuf := make([]byte, 64)
|
|
|
|
|
bBuf := make([]byte, 64)
|
|
|
|
|
cBuf := make([]byte, 64)
|
|
|
|
|
want := make([]float32, 16)
|
|
|
|
|
for i := range 16 {
|
|
|
|
|
p := pairs[i%len(pairs)]
|
|
|
|
|
ah, bh := bf16Bits(p[0]), bf16Bits(p[1])
|
|
|
|
|
al, bl := bf16Bits(p[2]), bf16Bits(p[3])
|
|
|
|
|
binary.LittleEndian.PutUint32(aBuf[4*i:], uint32(ah)<<16|uint32(al))
|
|
|
|
|
binary.LittleEndian.PutUint32(bBuf[4*i:], uint32(bh)<<16|uint32(bl))
|
|
|
|
|
want[i] = 1.0 + p[0]*p[1] + p[2]*p[3]
|
|
|
|
|
binary.LittleEndian.PutUint32(cBuf[4*i:], math.Float32bits(1.0))
|
|
|
|
|
}
|
|
|
|
|
args := make([]byte, 24)
|
|
|
|
|
PutPtr(args, 0, unsafe.Pointer(&aBuf[0]))
|
|
|
|
|
PutPtr(args, 8, unsafe.Pointer(&bBuf[0]))
|
|
|
|
|
PutPtr(args, 16, unsafe.Pointer(&cBuf[0]))
|
|
|
|
|
if _, err := k.CallFunc("dp", args); err != nil {
|
|
|
|
|
t.Fatalf("CallFunc: %v", err)
|
|
|
|
|
}
|
|
|
|
|
for i := range 16 {
|
|
|
|
|
if got := math.Float32frombits(binary.LittleEndian.Uint32(cBuf[4*i:])); got != want[i] {
|
|
|
|
|
t.Errorf("lane %d accumulated %v, want %v", i, got, want[i])
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// bf16Bits rounds a float32 value into its BF16 encoding.
|
|
|
|
|
func bf16Bits(f float32) uint16 {
|
|
|
|
|
return bf16Round(math.Float32bits(f))
|
|
|
|
|
}
|