test(verify): accumulate the BF16 dot product on the metal

Assisted-by: GLM 5.3 Flash
This commit is contained in:
petrbalvin committed 2026-10-07 00:07:58 +02:00
1 parent 22055b9bf3
commit 127e52de69
1 file changed
+66
+66
View File
@@ -199,3 +199,69 @@ TEXT ·isect(SB), NOSPLIT, $0-24
t.Errorf("the intersection masks are %#018x, want %#018x", got, want)
}
}
// 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))
}