AVX512汇编函数在Go协程并发调用时异常问题排查
自定义AVX512汇编函数并发调用失败排查与解决
问题概述
使用Go 1.23.0编写的AVX512汇编函数SubsetAVX512,通过int64位集实现手牌校验逻辑,单线程测试正常,但多goroutine并发调用时测试失败,怀疑是寄存器被协程覆盖导致。
汇编代码(原始版本)
// func SubsetAVX512(cs []CardSet, hs []CardSet) int // Returns 1 if any card set in cards contains any hand in hands, 0 otherwise #include "textflag.h" #define cs_data 0(FP) #define cs_len 8(FP) #define cs_cap 16(FP) #define hs_data 24(FP) #define hs_len 32(FP) #define hs_cap 40(FP) #define ret_off 48(FP) // Define the function TEXT ·SubsetAVX512(SB), NOSPLIT, $0-56 // Start of the function // Load parameters into registers MOVQ cs+cs_data, R8 // R8 = cards_ptr MOVQ cs+cs_len, R9 // R9 = cards_len MOVQ hs+hs_data, R10 // R10 = hands_ptr MOVQ hs+hs_len, R11 // R11 = hands_len // Check if hands_len == 0 TESTQ R11, R11 JE return_false // Check if cards_len == 0 TESTQ R9, R9 JE return_false // Initialize loop counters XORQ R12, R12 // R12 = i = 0 (hands index) // Main loop over hands outer_loop: CMPQ R12, R11 // Compare i (R12) with hands_len (R11) JGE return_false // If i >= hands_len, no match found // Load 8 hands into Z0 (512-bit register) LEAQ (R10)(R12*8), R13 // R13 = &hands[i] VMOVDQU64 0(R13), Z0 // Load 8 int64s from [R13] into Z0 // Inner loop over cards XORQ R14, R14 // R14 = j = 0 (cards index) inner_loop: CMPQ R14, R9 // Compare j (R14) with cards_len (R9) JGE next_hands_block // If j >= cards_len, move to next hands block // Load cs from cards[j] LEAQ (R8)(R14*8), R15 // R15 = &cards[j] MOVQ 0(R15), AX // AX = cards[j] // Broadcast cs into Z1 VPBROADCASTQ AX, Z1 // Broadcast RAX into all lanes of Z1 // Compute cs_vec & h_vec VPANDQ Z0, Z1, Z2 // Z2 = Z0 & Z1 // Compare (cs_vec & h_vec) == h_vec VPCMPEQQ Z0, Z2, K1 // Compare Z0 == Z2, store result in mask K1 // Check if any comparison is true KORTESTW K1, K1 // Test if any bits in K1 are set JNZ found_match // If so, a match is found // Increment card index INCQ R14 // j++ JMP inner_loop // Repeat inner loop next_hands_block: // Increment hands index by 8 ADDQ $8, R12 // i += 8 JMP outer_loop // Repeat outer loop found_match: // Match found, return 1 MOVQ $1, AX // Set return value to 1 (true) RET return_false: // No match found, return 0 XORQ AX, AX // Set return value to 0 (false) RET
测试代码
单线程测试(正常运行)
type CardSet int64 func SubsetAVX512(cs, hs []CardSet) bool func TestSubsetAVX512(t *testing.T) { cs := []CardSet{3, 1} hs := []CardSet{3, 0} var count int64 for i := 0; i < 5; i++ { if SubsetAVX512(cs, hs) { atomic.AddInt64(&count, 1) } } require.Equal(t, int64(5), count) }
并发测试(运行失败)
type CardSet int64 func SubsetAVX512(cs, hs []CardSet) bool func TestSubsetAVX512(t *testing.T) { cs := []CardSet{3, 1} hs := []CardSet{3, 0} var count int64 wg := sync.WaitGroup{} for i := 0; i < 5; i++ { wg.Add(1) go func() { defer wg.Done() if SubsetAVX512(cs, hs) { atomic.AddInt64(&count, 1) } }() } wg.Wait() require.Equal(t, int64(5), count) }
问题原因
- 被调用者保存寄存器未正确保存:Go的AMD64调用约定中,
R12-R15属于被调用者保存寄存器,函数修改这些寄存器时必须先保存原值,执行完毕后恢复。原始代码直接修改了R12-R15但未做保存恢复,会破坏调用者(Go runtime)的上下文,并发调度时会引发错误。 - AVX512寄存器未处理调用者保存规则:AVX512的
ZMM寄存器(如Z0-Z2)和K掩码寄存器(如K1)属于调用者保存寄存器,当goroutine被抢占时,其他协程会修改这些寄存器的值,导致本函数恢复执行时拿到错误的寄存器内容,逻辑判断出错。
解决方法
修改汇编代码,在函数开头保存所有被修改的寄存器,执行完毕后恢复:
// func SubsetAVX512(cs []CardSet, hs []CardSet) int // Returns 1 if any card set in cards contains any hand in hands, 0 otherwise #include "textflag.h" #define cs_data 0(FP) #define cs_len 8(FP) #define cs_cap 16(FP) #define hs_data 24(FP) #define hs_len 32(FP) #define hs_cap 40(FP) #define ret_off 48(FP) // 栈空间:4个通用寄存器(32字节) + 3个ZMM寄存器(192字节) + 1个K寄存器(8字节) = 232字节,对齐到16字节为240字节 TEXT ·SubsetAVX512(SB), NOSPLIT, $240-56 // Start of the function // 保存被调用者保存的寄存器R12-R15 MOVQ R12, 0(SP) MOVQ R13, 8(SP) MOVQ R14, 16(SP) MOVQ R15, 24(SP) // 保存调用者保存的AVX512寄存器Z0-Z2 VMOVDQU64 Z0, 32(SP) VMOVDQU64 Z1, 96(SP) VMOVDQU64 Z2, 160(SP) // 保存掩码寄存器K1 KMOVD K1, AX MOVQ AX, 224(SP) // Load parameters into registers MOVQ cs+cs_data, R8 // R8 = cards_ptr MOVQ cs+cs_len, R9 // R9 = cards_len MOVQ hs+hs_data, R10 // R10 = hands_ptr MOVQ hs+hs_len, R11 // R11 = hands_len // Check if hands_len == 0 TESTQ R11, R11 JE return_false // Check if cards_len == 0 TESTQ R9, R9 JE return_false // Initialize loop counters XORQ R12, R12 // R12 = i = 0 (hands index) // Main loop over hands outer_loop: CMPQ R12, R11 // Compare i (R12) with hands_len (R11) JGE return_false // If i >= hands_len, no match found // Load 8 hands into Z0 (512-bit register) LEAQ (R10)(R12*8), R13 // R13 = &hands[i] VMOVDQU64 0(R13), Z0 // Load 8 int64s from [R13] into Z0 // Inner loop over cards XORQ R14, R14 // R14 = j = 0 (cards index) inner_loop: CMPQ R14, R9 // Compare j (R14) with cards_len (R9) JGE next_hands_block // If j >= cards_len, move to next hands block // Load cs from cards[j] LEAQ (R8)(R14*8), R15 // R15 = &cards[j] MOVQ 0(R15), AX // AX = cards[j] // Broadcast cs into Z1 VPBROADCASTQ AX, Z1 // Broadcast RAX into all lanes of Z1 // Compute cs_vec & h_vec VPANDQ Z0, Z1, Z2 // Z2 = Z0 & Z1 // Compare (cs_vec & h_vec) == h_vec VPCMPEQQ Z0, Z2, K1 // Compare Z0 == Z2, store result in mask K1 // Check if any comparison is true KORTESTW K1, K1 // Test if any bits in K1 are set JNZ found_match // If so, a match is found // Increment card index INCQ R14 // j++ JMP inner_loop // Repeat inner loop next_hands_block: // Increment hands index by 8 ADDQ $8, R12 // i += 8 JMP outer_loop // Repeat outer loop found_match: // Match found, return 1 MOVQ $1, AX // Set return value to 1 (true) // 跳转到寄存器恢复逻辑 JMP restore_regs return_false: // No match found, return 0 XORQ AX, AX // Set return value to 0 (false) restore_regs: // 恢复掩码寄存器K1 MOVQ 224(SP), AX KMOVD AX, K1 // 恢复AVX512寄存器Z0-Z2 VMOVDQU64 32(SP), Z0 VMOVDQU64 96(SP), Z1 VMOVDQU64 160(SP), Z2 // 恢复通用寄存器R12-R15 MOVQ 0(SP), R12 MOVQ 8(SP), R13 MOVQ 16(SP), R14 MOVQ 24(SP), R15 ADDQ $240, SP RET
验证说明
修改后的代码在函数入口保存了所有会被修改的寄存器,在返回前恢复原值,符合Go的调用约定,并发调用时不会再出现寄存器被覆盖的问题,并发测试可正常通过。
内容的提问来源于stack exchange,提问作者Joe Doliner
相关产品推荐
相关产品推荐

