diff --git a/verify/amd64_ext_jit_test.go b/verify/amd64_ext_jit_test.go index 1a05dd0..bdeef22 100644 --- a/verify/amd64_ext_jit_test.go +++ b/verify/amd64_ext_jit_test.go @@ -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)) +}