test(verify): accumulate the BF16 dot product on the metal
Assisted-by: GLM 5.3 Flash
This commit is contained in:
1 parent
22055b9bf3
commit
127e52de69
1 file changed
+66
@@ -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))
|
||||
}
|
||||
Reference in new issue
Block a user