From 36a7120662d11918ddecfba04fc657f7f317e04e Mon Sep 17 00:00:00 2001 From: Lancelot Six Date: Fri, 16 Aug 2024 19:49:58 +0000 Subject: [PATCH] rocr/trap_handler_gfx12: Properly ignore HOST_TRAP and debug_trap The current trap handler has 2 limitations: 1) If it receives a HOST_TRAP, it clears the corresponding bit and notifies the host, when it should not. 2) When it is entered because of a debug trap (s_trap 3) and the debugger is not attached, it returns unconditionally. However, if another exception is reported at the same time as the trap handler is entered for the debug trap (a memory violation for example), that other exception ends-up being ignored. This patch addresses both of those issues. It makes it so host traps and debug traps are ignored when necessary. If any other exception is reported to the wave, we halt the wave and notify the host, and if no other exception is reported (i.e. we entered the trap handler because of host trap or debug trap), we return to shader code. Other minor defects are also fixed during this refactor: - Fixed SQ_WAVE_EXCP_FLAG_PRIV_XNACK_ERROR_SHIFT which had an incorrect value - Host traps can be sent at any time, including after we have halted a wave. In such case, the old approach would have: 1) cleared the trap ID saved in ttmp6 2) clobbered ttmp10 where part of the actual wave's PC is saved. Change-Id: I9ecd341f4967e686233dec182b3e5b0388ef19bd [ROCm/ROCR-Runtime commit: 123b2c080a3505b5da32d8c94951c5f9be21672e] --- .../runtime/trap_handler/trap_handler_gfx12.s | 133 ++++++++++++------ 1 file changed, 93 insertions(+), 40 deletions(-) diff --git a/projects/rocr-runtime/runtime/hsa-runtime/core/runtime/trap_handler/trap_handler_gfx12.s b/projects/rocr-runtime/runtime/hsa-runtime/core/runtime/trap_handler/trap_handler_gfx12.s index 370cd2af28..f7fd15610f 100644 --- a/projects/rocr-runtime/runtime/hsa-runtime/core/runtime/trap_handler/trap_handler_gfx12.s +++ b/projects/rocr-runtime/runtime/hsa-runtime/core/runtime/trap_handler/trap_handler_gfx12.s @@ -49,14 +49,24 @@ .set EC_QUEUE_WAVE_ILLEGAL_INSTRUCTION_M0 , (1 << (DOORBELL_ID_SIZE + 3)) .set EC_QUEUE_WAVE_MEMORY_VIOLATION_M0 , (1 << (DOORBELL_ID_SIZE + 4)) .set EC_QUEUE_WAVE_APERTURE_VIOLATION_M0 , (1 << (DOORBELL_ID_SIZE + 5)) + +.set SQ_WAVE_EXCP_FLAG_PRIV_ADDR_WATCH_MASK , (1 << 4) - 1 .set SQ_WAVE_EXCP_FLAG_PRIV_MEMVIOL_SHIFT , 4 -.set SQ_WAVE_EXCP_FLAG_PRIV_HT_SHIFT , 7 .set SQ_WAVE_EXCP_FLAG_PRIV_ILLEGAL_INST_SHIFT , 6 -.set SQ_WAVE_EXCP_FLAG_PRIV_XNACK_ERROR_SHIFT , 8 +.set SQ_WAVE_EXCP_FLAG_PRIV_HT_SHIFT , 7 +.set SQ_WAVE_EXCP_FLAG_PRIV_WAVE_START_SHIFT , 8 +.set SQ_WAVE_EXCP_FLAG_PRIV_WAVE_END_SHIFT , 9 +.set SQ_WAVE_EXCP_FLAG_PRIV_TRAP_AFTER_INST_SHIFT , 11 +.set SQ_WAVE_EXCP_FLAG_PRIV_XNACK_ERROR_SHIFT , 12 + .set SQ_WAVE_EXCP_FLAG_USER_MATH_EXCP_SHIFT , 0 .set SQ_WAVE_EXCP_FLAG_USER_MATH_EXCP_SIZE , 7 -.set SQ_WAVE_TRAP_CTRL_MATH_EXCP_SHIFT , 0 -.set SQ_WAVE_TRAP_CTRL_MATH_EXCP_SIZE , 7 + +.set SQ_WAVE_TRAP_CTRL_MATH_EXCP_MASK , ((1 << 7) - 1) +.set SQ_WAVE_TRAP_CTRL_ADDR_WATCH_SHIFT , 7 +.set SQ_WAVE_TRAP_CTRL_WAVE_END_SHIFT , 8 +.set SQ_WAVE_TRAP_CTRL_TRAP_AFTER_INST , 9 + .set SQ_WAVE_PC_HI_ADDRESS_MASK , 0xFFFF .set SQ_WAVE_PC_HI_TRAP_ID_BFE , (SQ_WAVE_PC_HI_TRAP_ID_SHIFT | (SQ_WAVE_PC_HI_TRAP_ID_SIZE << 16)) .set SQ_WAVE_PC_HI_TRAP_ID_SHIFT , 28 @@ -85,41 +95,52 @@ // ttmp15 = TMA[63:32] trap_entry: + // Clear ttmp3 as it will contain the exception code. + s_mov_b32 ttmp3, 0 + // Branch if not a trap (an exception instead). s_bfe_u32 ttmp2, ttmp1, SQ_WAVE_PC_HI_TRAP_ID_BFE - s_cbranch_scc0 .no_skip_debugtrap + s_cbranch_scc0 .check_exceptions - // If caused by s_trap then advance PC. + // If caused by s_trap then advance PC, then figure out the trap ID: + // - if trapID is DEBUGTRAP and debugger is attach, report WAVE_TRAP, + // - if trapID is ABORTTRAP, report WAVE_ABORT, + // - report WAVE_TRAP for any other trap ID. s_add_u32 ttmp0, ttmp0, 0x4 s_addc_u32 ttmp1, ttmp1, 0x0 -.not_s_trap: // If llvm.debugtrap and debugger is not attached. s_cmp_eq_u32 ttmp2, TRAP_ID_DEBUGTRAP - s_cbranch_scc0 .no_skip_debugtrap + s_cbranch_scc0 .not_debug_trap - s_bitcmp0_b32 ttmp11, TTMP11_DEBUG_ENABLED_SHIFT - s_cbranch_scc0 .no_skip_debugtrap + s_bitcmp1_b32 ttmp11, TTMP11_DEBUG_ENABLED_SHIFT + s_cbranch_scc0 .check_exceptions + s_or_b32 ttmp3, ttmp3, EC_QUEUE_WAVE_TRAP_M0 - // Ignore llvm.debugtrap. - s_branch .exit_trap +.not_debug_trap: + s_cmp_eq_u32 ttmp2, TRAP_ID_ABORT + s_cbranch_scc0 .not_abort_trap + s_or_b32 ttmp3, ttmp3, EC_QUEUE_WAVE_ABORT_M0 + s_branch .check_exceptions -.no_skip_debugtrap: - // Save trap id and halt status in ttmp6. - s_andn2_b32 ttmp6, ttmp6, (TTMP6_SAVED_TRAP_ID_MASK | TTMP6_SAVED_STATUS_HALT_MASK) - s_min_u32 ttmp2, ttmp2, 0xF - s_lshl_b32 ttmp2, ttmp2, TTMP6_SAVED_TRAP_ID_SHIFT - s_or_b32 ttmp6, ttmp6, ttmp2 - s_bfe_u32 ttmp2, ttmp12, SQ_WAVE_STATE_PRIV_HALT_BFE - s_lshl_b32 ttmp2, ttmp2, TTMP6_SAVED_STATUS_HALT_SHIFT - s_or_b32 ttmp6, ttmp6, ttmp2 - - // Fetch doorbell id for our queue. - s_sendmsg_rtn_b32 ttmp3, sendmsg(MSG_RTN_GET_DOORBELL) - s_wait_kmcnt 0 - s_and_b32 ttmp3, ttmp3, DOORBELL_ID_MASK +.not_abort_trap: + s_or_b32 ttmp3, ttmp3, EC_QUEUE_WAVE_TRAP_M0 + // We need to explititly look for all exceptions we want to report to the + // host: + // - EXCP_FLAG_PRIV.XNACK_ERROR (&& EXCP_FLAG_PRIV.MEMVIOL) + // -> WAVE_MEMORY_VIOLATION + // - EXCP_FLAG_PRIV.MEMVIOL (and !EXCP_FLAG_PRIV.XNACK_ERROR) + // -> WAVE_APERTURE_VIOLATION + // - EXCP_FLAG_PRIV.ILLEGAL_INST -> WAVE_ILLEGAL_INSTRUCTION + // - EXCP_FLAG_PRIV.WAVE_START -> WAVE_TRAP + // - EXCP_FLAG_PRIV.WAVE_END && TRAP_CTRL.WAVE_END -> WAVE_TRAP + // - TRAP_CTRL.TRAP_AFTER_INST -> WAVE_TRAP + // - EXCP_FLAG_PRIV.ADDR_WATCH && TRAP_CTL.WATCH -> WAVE_TRAP + // - (EXCP_FLAG_USER[ALU] & TRAP_CTRL[ALU]) != 0 -> WAVE_MATH_ERROR +.check_exceptions: s_getreg_b32 ttmp2, hwreg(HW_REG_EXCP_FLAG_PRIV) + s_getreg_b32 ttmp13, hwreg(HW_REG_TRAP_CTRL) s_bitcmp1_b32 ttmp2, SQ_WAVE_EXCP_FLAG_PRIV_XNACK_ERROR_SHIFT s_cbranch_scc0 .not_memory_violation @@ -139,27 +160,59 @@ trap_entry: s_or_b32 ttmp3, ttmp3, EC_QUEUE_WAVE_ILLEGAL_INSTRUCTION_M0 .not_illegal_instruction: - s_getreg_b32 ttmp2, hwreg(HW_REG_EXCP_FLAG_USER, SQ_WAVE_EXCP_FLAG_USER_MATH_EXCP_SHIFT, SQ_WAVE_EXCP_FLAG_USER_MATH_EXCP_SIZE) - s_cbranch_scc0 .not_math_exception - s_getreg_b32 ttmp10, hwreg(HW_REG_TRAP_CTRL, SQ_WAVE_TRAP_CTRL_MATH_EXCP_SHIFT, SQ_WAVE_TRAP_CTRL_MATH_EXCP_SIZE) - s_and_b32 ttmp2, ttmp2, ttmp10 + s_bitcmp1_b32 ttmp2, SQ_WAVE_EXCP_FLAG_PRIV_WAVE_START_SHIFT + s_cbranch_scc0 .not_wave_end + s_or_b32 ttmp3, ttmp3, EC_QUEUE_WAVE_TRAP_M0 +.not_wave_start: + s_bitcmp1_b32 ttmp2, SQ_WAVE_EXCP_FLAG_PRIV_WAVE_END_SHIFT + s_cbranch_scc0 .not_wave_end + s_bitcmp1_b32 ttmp13, SQ_WAVE_TRAP_CTRL_WAVE_END_SHIFT + s_cbranch_scc0 .not_wave_end + s_or_b32 ttmp3, ttmp3, EC_QUEUE_WAVE_TRAP_M0 + +.not_wave_end: + s_bitcmp1_b32 ttmp13, SQ_WAVE_TRAP_CTRL_TRAP_AFTER_INST + s_cbranch_scc0 .not_trap_after_inst + s_or_b32 ttmp3, ttmp3, EC_QUEUE_WAVE_TRAP_M0 + +.not_trap_after_inst: + s_and_b32 ttmp2, ttmp2, SQ_WAVE_EXCP_FLAG_PRIV_ADDR_WATCH_MASK + s_cbranch_scc0 .not_addr_watch + s_bitcmp1_b32 ttmp13, SQ_WAVE_TRAP_CTRL_ADDR_WATCH_SHIFT + s_cbranch_scc0 .not_addr_watch + s_or_b32 ttmp3, ttmp3, EC_QUEUE_WAVE_TRAP_M0 + +.not_addr_watch: + s_getreg_b32 ttmp2, hwreg(HW_REG_EXCP_FLAG_USER, SQ_WAVE_EXCP_FLAG_USER_MATH_EXCP_SHIFT, SQ_WAVE_EXCP_FLAG_USER_MATH_EXCP_SIZE) + s_and_b32 ttmp13, ttmp13, SQ_WAVE_TRAP_CTRL_MATH_EXCP_MASK + s_and_b32 ttmp2, ttmp2, ttmp13 s_cbranch_scc0 .not_math_exception s_or_b32 ttmp3, ttmp3, EC_QUEUE_WAVE_MATH_ERROR_M0 .not_math_exception: - s_bfe_u32 ttmp2, ttmp6, TTMP6_SAVED_TRAP_ID_BFE - s_cmp_eq_u32 ttmp2, TRAP_ID_ABORT - s_cbranch_scc0 .not_abort_trap - s_or_b32 ttmp3, ttmp3, EC_QUEUE_WAVE_ABORT_M0 - -.not_abort_trap: - // If no other exception was flagged then report a generic error. - s_andn2_b32 ttmp2, ttmp3, DOORBELL_ID_MASK - s_cbranch_scc1 .send_interrupt - s_or_b32 ttmp3, ttmp3, EC_QUEUE_WAVE_TRAP_M0 + s_cmp_eq_u32 ttmp3, 0 + // This was not a s_trap we are interested in or an exception, return to + // the user code. + s_cbranch_scc1 .exit_trap .send_interrupt: + // Fetch doorbell id for our queue. + s_sendmsg_rtn_b32 ttmp2, sendmsg(MSG_RTN_GET_DOORBELL) + s_wait_kmcnt 0 + s_and_b32 ttmp2, ttmp2, DOORBELL_ID_MASK + s_or_b32 ttmp3, ttmp2, ttmp3 + + // Save trap id and halt status in ttmp6. + s_andn2_b32 ttmp6, ttmp6, (TTMP6_SAVED_TRAP_ID_MASK | TTMP6_SAVED_STATUS_HALT_MASK) + s_bfe_u32 ttmp2, ttmp1, SQ_WAVE_PC_HI_TRAP_ID_BFE + s_min_u32 ttmp2, ttmp2, 0xF + s_lshl_b32 ttmp2, ttmp2, TTMP6_SAVED_TRAP_ID_SHIFT + s_or_b32 ttmp6, ttmp6, ttmp2 + s_bfe_u32 ttmp2, ttmp12, SQ_WAVE_STATE_PRIV_HALT_BFE + s_lshl_b32 ttmp2, ttmp2, TTMP6_SAVED_STATUS_HALT_SHIFT + s_or_b32 ttmp6, ttmp6, ttmp2 + // m0 = interrupt data = (exception_code << DOORBELL_ID_SIZE) | doorbell_id s_mov_b32 ttmp2, m0 s_mov_b32 m0, ttmp3