From c0fbe34efc79e0ea56828975794a878f242347c7 Mon Sep 17 00:00:00 2001 From: Zhendong404 Date: Wed, 5 Aug 2026 00:31:07 +0800 Subject: [PATCH 1/2] Fix VPTO trap LLVM lowering --- lib/PTO/Transforms/VPTOCANN900LLVMEmitter.cpp | 30 +++++++++++++++++-- lib/PTO/Transforms/VPTOLLVMEmitter.cpp | 30 +++++++++++++++++-- test/lit/vpto/trap_vpto_llvm.pto | 24 +++++++++++++++ 3 files changed, 80 insertions(+), 4 deletions(-) create mode 100644 test/lit/vpto/trap_vpto_llvm.pto diff --git a/lib/PTO/Transforms/VPTOCANN900LLVMEmitter.cpp b/lib/PTO/Transforms/VPTOCANN900LLVMEmitter.cpp index 9c50f1ec7e..bde7bb4faa 100644 --- a/lib/PTO/Transforms/VPTOCANN900LLVMEmitter.cpp +++ b/lib/PTO/Transforms/VPTOCANN900LLVMEmitter.cpp @@ -298,6 +298,29 @@ struct LoweringState { SmallVector plannedDecls; }; +class LowerTrapOpPattern final : public OpConversionPattern { +public: + explicit LowerTrapOpPattern(TypeConverter &typeConverter, + MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), + state(state) {} + + LogicalResult + matchAndRewrite(pto::TrapOp op, pto::TrapOp::Adaptor, + ConversionPatternRewriter &rewriter) const override { + constexpr StringLiteral calleeName = "llvm.hivm.TRAP"; + auto funcType = rewriter.getFunctionType({}, {}); + rewriter.create(op.getLoc(), calleeName, TypeRange{}, + ValueRange{}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + enum class VcvtElemKind { Invalid, F16, @@ -10611,6 +10634,7 @@ static void populateVPTOOpLoweringPatterns(VPTOTypeConverter &typeConverter, LowerAtomicBinaryOpPattern, LowerAtomicBinaryOpPattern, LowerAtomicBinaryOpPattern, + LowerTrapOpPattern, LowerScalarIntrinsicOpPattern, LowerMulhiOpPattern, LowerMulI32ToI64OpPattern, @@ -10753,7 +10777,7 @@ static void configureVPTOOpLoweringTarget(ConversionTarget &target, pto::AtomicAddOp, pto::AtomicSubOp, pto::AtomicMinOp, pto::AtomicMaxOp, pto::AtomicAndOp, pto::AtomicOrOp, - pto::AtomicXorOp, pto::PrmtOp, + pto::AtomicXorOp, pto::TrapOp, pto::PrmtOp, pto::MulhiOp, pto::MulI32ToI64Op, pto::SqrtOp, pto::AbsFOp, pto::ExpOp, pto::LogOp, pto::CeilOp, pto::FloorOp, pto::RintOp, pto::RoundOp, pto::FMinOp, @@ -10828,7 +10852,9 @@ static void configureVPTOOpLoweringTarget(ConversionTarget &target, pto::MadMxAccOp, pto::MadMxBiasOp, pto::MadRawOp, pto::MadBiasRawOp, pto::MadMxRawOp, pto::MadMxBiasRawOp>(); - target.markUnknownOpDynamicallyLegal([](Operation *) { return true; }); + target.markUnknownOpDynamicallyLegal([](Operation *op) { + return !isa(op); + }); } static void populateVPTOStructuralTypePatterns( diff --git a/lib/PTO/Transforms/VPTOLLVMEmitter.cpp b/lib/PTO/Transforms/VPTOLLVMEmitter.cpp index e7d5a79906..64daa20d47 100644 --- a/lib/PTO/Transforms/VPTOLLVMEmitter.cpp +++ b/lib/PTO/Transforms/VPTOLLVMEmitter.cpp @@ -300,6 +300,29 @@ struct LoweringState { SmallVector plannedDecls; }; +class LowerTrapOpPattern final : public OpConversionPattern { +public: + explicit LowerTrapOpPattern(TypeConverter &typeConverter, + MLIRContext *context, LoweringState &state) + : OpConversionPattern(typeConverter, context), + state(state) {} + + LogicalResult + matchAndRewrite(pto::TrapOp op, pto::TrapOp::Adaptor, + ConversionPatternRewriter &rewriter) const override { + constexpr StringLiteral calleeName = "llvm.hivm.TRAP"; + auto funcType = rewriter.getFunctionType({}, {}); + rewriter.create(op.getLoc(), calleeName, TypeRange{}, + ValueRange{}); + state.plannedDecls.push_back(PlannedDecl{calleeName.str(), funcType}); + rewriter.eraseOp(op); + return success(); + } + +private: + LoweringState &state; +}; + enum class VcvtElemKind { Invalid, F16, @@ -11266,6 +11289,7 @@ static void populateVPTOOpLoweringPatterns(VPTOTypeConverter &typeConverter, LowerAtomicBinaryOpPattern, LowerAtomicBinaryOpPattern, LowerAtomicBinaryOpPattern, + LowerTrapOpPattern, LowerScalarIntrinsicOpPattern, LowerMulhiOpPattern, LowerMulI32ToI64OpPattern, @@ -11474,7 +11498,7 @@ static void configureVPTOOpLoweringTarget(ConversionTarget &target, pto::AtomicAddOp, pto::AtomicSubOp, pto::AtomicMinOp, pto::AtomicMaxOp, pto::AtomicAndOp, pto::AtomicOrOp, - pto::AtomicXorOp, pto::PrmtOp, + pto::AtomicXorOp, pto::TrapOp, pto::PrmtOp, pto::MulhiOp, pto::MulI32ToI64Op, pto::SqrtOp, pto::AbsFOp, pto::ExpOp, pto::LogOp, pto::CeilOp, pto::FloorOp, pto::RintOp, pto::RoundOp, pto::FMinOp, @@ -11579,7 +11603,9 @@ static void configureVPTOOpLoweringTarget(ConversionTarget &target, target.addIllegalOp(); } - target.markUnknownOpDynamicallyLegal([](Operation *) { return true; }); + target.markUnknownOpDynamicallyLegal([](Operation *op) { + return !isa(op); + }); } static void populateVPTOStructuralTypePatterns( diff --git a/test/lit/vpto/trap_vpto_llvm.pto b/test/lit/vpto/trap_vpto_llvm.pto new file mode 100644 index 0000000000..902b440462 --- /dev/null +++ b/test/lit/vpto/trap_vpto_llvm.pto @@ -0,0 +1,24 @@ +// Copyright (c) 2026 Huawei Technologies Co., Ltd. +// This program is free software, you can redistribute it and/or modify it under the terms of +// CANN Open Software License Agreement Version 2.0 (the "License"). +// Please refer to the License for details in the License. +// THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER IMPLIED, +// INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY OR FITNESS FOR A PARTICULAR PURPOSE. +// See the License for the full text of the License. + +// pto.trap must lower to Bisheng's AICore trap intrinsic before LLVM export. +// RUN: ptoas --pto-arch=a5 --pto-backend=vpto --emit-vpto-llvm-ir %s -o - 2>&1 | FileCheck %s +// RUN: ptoas --cann-output-version=9.0.0 --pto-arch=a5 --pto-backend=vpto --emit-vpto-llvm-ir %s -o - 2>&1 | FileCheck %s + +module attributes {pto.target_arch = "a5", pto.kernel_kind = #pto.kernel_kind} { + func.func @trap_vpto_llvm(%arg0: i1) attributes {pto.entry} { + scf.if %arg0 { + pto.trap + } + return + } +} + +// CHECK-DAG: declare void @llvm.hivm.TRAP() +// CHECK-LABEL: define void @trap_vpto_llvm_mix_aiv +// CHECK: call void @llvm.hivm.TRAP() From fdfd645c0fe0cacd431fcf2e83725df44b1bf861 Mon Sep 17 00:00:00 2001 From: Zhendong404 Date: Thu, 6 Aug 2026 13:56:04 +0800 Subject: [PATCH 2/2] Expose device trap in PTODSL --- ptodsl/docs/user_guide/10-sync-ops.md | 21 +++++++++++++++++++ ptodsl/ptodsl/_ops.py | 7 ++++++- ptodsl/ptodsl/pto.py | 2 +- .../tests/support/docs_fragment_fixtures.py | 7 +++++++ ptodsl/tests/test_jit_compile.py | 9 ++++++++ 5 files changed, 44 insertions(+), 2 deletions(-) diff --git a/ptodsl/docs/user_guide/10-sync-ops.md b/ptodsl/docs/user_guide/10-sync-ops.md index dcd6917ef8..e1fed249b2 100644 --- a/ptodsl/docs/user_guide/10-sync-ops.md +++ b/ptodsl/docs/user_guide/10-sync-ops.md @@ -489,3 +489,24 @@ In auto mode, users can still write sync operations directly — `set_flag`/`wai | Double-buffer handoff (compute → DMA) | `rls_buf(V, id)` + `get_buf(MTE2, id)` | | Double-buffer handoff (DMA → compute) | `rls_buf(MTE2, id)` + `get_buf(V, id)` | | Core A notifies core B | `set_cross_flag(B, id)` + `wait_cross_flag(A, id)` | + +## 10.7 Device-side trap + +### `pto.trap()` + +**Description**: Unconditionally terminates execution of the current device-side +execution instance. This is useful for fail-fast paths in lowered assertions or +debug-only guards. + +`pto.trap()` is a device operation emitted while the kernel is being traced. It +is not a Python exception and must not be replaced with Python `raise` or +`assert`, which execute during host-side tracing instead of on the device. + +**Returns**: None. Execution does not continue after the trap at runtime. + + +```python +# The condition and any diagnostic reporting are normally emitted by a +# higher-level assertion helper; trap itself is unconditional. +pto.trap() +``` diff --git a/ptodsl/ptodsl/_ops.py b/ptodsl/ptodsl/_ops.py index 365ec64b04..1436a640cb 100644 --- a/ptodsl/ptodsl/_ops.py +++ b/ptodsl/ptodsl/_ops.py @@ -6210,6 +6210,11 @@ def threadfence_block(): _pto.ThreadfenceBlockOp() +def trap(): + """``pto.trap`` – unconditionally terminate device-side execution.""" + _pto.TrapOp() + + def _slot_attr_value(slot, *, context: str): if not isinstance(slot, int) or isinstance(slot, bool): raise TypeError(f"{context} expects a non-negative Python int slot") @@ -6492,7 +6497,7 @@ def import_reserved_buffer(name, *, peer_func): "prmt", "mulhi", "mul_i32toi64", "absf", "sqrt", "exp", "log", "pow", "ceil", "floor", "rint", "round", "fmin", "fmax", "fma", "convert", - "syncthreads", "threadfence", "threadfence_block", "keep", "resume", + "syncthreads", "threadfence", "threadfence_block", "trap", "keep", "resume", "pipe_barrier", "get_buf", "rls_buf", "set_cross_flag", "wait_cross_flag", "set_intra_flag", "wait_intra_flag", "set_flag", "wait_flag", diff --git a/ptodsl/ptodsl/pto.py b/ptodsl/ptodsl/pto.py index c45d521700..7fc6f426db 100644 --- a/ptodsl/ptodsl/pto.py +++ b/ptodsl/ptodsl/pto.py @@ -148,7 +148,7 @@ prmt, mulhi, mul_i32toi64, absf, sqrt, exp, log, pow, ceil, floor, rint, round, fmin, fmax, fma, convert, - syncthreads, threadfence, threadfence_block, keep, resume, + syncthreads, threadfence, threadfence_block, trap, keep, resume, pipe_barrier, get_buf, rls_buf, set_cross_flag, wait_cross_flag, set_intra_flag, wait_intra_flag, diff --git a/ptodsl/tests/support/docs_fragment_fixtures.py b/ptodsl/tests/support/docs_fragment_fixtures.py index c773f046dc..fcd185e847 100644 --- a/ptodsl/tests/support/docs_fragment_fixtures.py +++ b/ptodsl/tests/support/docs_fragment_fixtures.py @@ -1354,6 +1354,13 @@ def sync_ops_basic_probe(): {SNIPPET_PLACEHOLDER} """ ), + "sync_ops.trap": _fixture( + f""" + @pto.jit(target="a5") + def sync_ops_trap_probe(): + {SNIPPET_PLACEHOLDER} + """ + ), "flash_attention.l1_tensor_views": _fixture( f""" @pto.jit(target="a5") diff --git a/ptodsl/tests/test_jit_compile.py b/ptodsl/tests/test_jit_compile.py index f40c2ed6f5..622bc1f313 100644 --- a/ptodsl/tests/test_jit_compile.py +++ b/ptodsl/tests/test_jit_compile.py @@ -3067,6 +3067,11 @@ def public_sync_surface_probe(): pto.wait_intra_flag(pto.Pipe.MTE3, 31) +@pto.jit(target="a5") +def public_trap_surface_probe(): + pto.trap() + + @pto.jit(target="a5") def public_dynamic_buf_sync_surface_probe(): const_buf_id = pto.const(3) @@ -3860,6 +3865,7 @@ def main() -> None: public_mask_bitcast_probe.verify() public_mask_surface_probe.verify() public_sync_surface_probe.verify() + public_trap_surface_probe.verify() explicit_runtime_index_bitwise_event_probe.verify() explicit_runtime_index_integer_bitwise_event_probe.verify() ast_runtime_index_bitwise_event_probe.verify() @@ -6274,6 +6280,8 @@ def _enter_inline_simt_with_resource_attr(): expect_parse_roundtrip_and_verify(mask_surface_text, "public mask surface specialization") sync_surface_text = public_sync_surface_probe.compile().mlir_text() expect_parse_roundtrip_and_verify(sync_surface_text, "public sync surface specialization") + trap_surface_text = public_trap_surface_probe.compile().mlir_text() + expect_parse_roundtrip_and_verify(trap_surface_text, "public trap surface specialization") dynamic_buf_sync_text = public_dynamic_buf_sync_surface_probe.compile().mlir_text() expect_parse_roundtrip_and_verify(dynamic_buf_sync_text, "dynamic buf sync surface specialization") explicit_runtime_index_bitwise_event_text = explicit_runtime_index_bitwise_event_probe.compile().mlir_text() @@ -6505,6 +6513,7 @@ def _enter_inline_simt_with_resource_attr(): ) expect(public_surface_text.count("pto.mem_bar") >= 1, "mem_bar(...) should still lower explicit memory barriers") expect("pto.barrier " in public_surface_text, "pipe_barrier(Pipe.ALL) should lower to pto.barrier") + expect("pto.trap" in trap_surface_text, "pto.trap() should lower to the public pto.trap operation") expect("pto.vexp" in public_surface_text, "vexp(...) should lower to pto.vexp") expect("pto.vcgmax" in public_surface_text, "vcgmax(...) should lower to pto.vcgmax") expect("pto.vcgadd" in public_surface_text, "vcgadd(...) should lower to pto.vcgadd")