Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 28 additions & 2 deletions lib/PTO/Transforms/VPTOCANN900LLVMEmitter.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -298,6 +298,29 @@ struct LoweringState {
SmallVector<PlannedDecl> plannedDecls;
};

class LowerTrapOpPattern final : public OpConversionPattern<pto::TrapOp> {
public:
explicit LowerTrapOpPattern(TypeConverter &typeConverter,
MLIRContext *context, LoweringState &state)
: OpConversionPattern<pto::TrapOp>(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<func::CallOp>(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,
Expand Down Expand Up @@ -10611,6 +10634,7 @@ static void populateVPTOOpLoweringPatterns(VPTOTypeConverter &typeConverter,
LowerAtomicBinaryOpPattern<pto::AtomicAndOp>,
LowerAtomicBinaryOpPattern<pto::AtomicOrOp>,
LowerAtomicBinaryOpPattern<pto::AtomicXorOp>,
LowerTrapOpPattern,
LowerScalarIntrinsicOpPattern<pto::PrmtOp>,
LowerMulhiOpPattern,
LowerMulI32ToI64OpPattern,
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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<pto::TrapOp>(op);
});
}

static void populateVPTOStructuralTypePatterns(
Expand Down
30 changes: 28 additions & 2 deletions lib/PTO/Transforms/VPTOLLVMEmitter.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -300,6 +300,29 @@ struct LoweringState {
SmallVector<PlannedDecl> plannedDecls;
};

class LowerTrapOpPattern final : public OpConversionPattern<pto::TrapOp> {
public:
explicit LowerTrapOpPattern(TypeConverter &typeConverter,
MLIRContext *context, LoweringState &state)
: OpConversionPattern<pto::TrapOp>(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<func::CallOp>(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,
Expand Down Expand Up @@ -11266,6 +11289,7 @@ static void populateVPTOOpLoweringPatterns(VPTOTypeConverter &typeConverter,
LowerAtomicBinaryOpPattern<pto::AtomicAndOp>,
LowerAtomicBinaryOpPattern<pto::AtomicOrOp>,
LowerAtomicBinaryOpPattern<pto::AtomicXorOp>,
LowerTrapOpPattern,
LowerScalarIntrinsicOpPattern<pto::PrmtOp>,
LowerMulhiOpPattern,
LowerMulI32ToI64OpPattern,
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -11579,7 +11603,9 @@ static void configureVPTOOpLoweringTarget(ConversionTarget &target,
target.addIllegalOp<pto::UBSetMaskNormOp>();
}

target.markUnknownOpDynamicallyLegal([](Operation *) { return true; });
target.markUnknownOpDynamicallyLegal([](Operation *op) {
return !isa<pto::TrapOp>(op);
});
}

static void populateVPTOStructuralTypePatterns(
Expand Down
21 changes: 21 additions & 0 deletions ptodsl/docs/user_guide/10-sync-ops.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.

<!-- ptodsl-doc-test: {"mode":"compile_fragment","fixture":"sync_ops.trap","symbol":"sync_ops_trap_probe","compile":{}} -->
```python
# The condition and any diagnostic reporting are normally emitted by a
# higher-level assertion helper; trap itself is unconditional.
pto.trap()
```
7 changes: 6 additions & 1 deletion ptodsl/ptodsl/_ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down Expand Up @@ -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",
Expand Down
2 changes: 1 addition & 1 deletion ptodsl/ptodsl/pto.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
7 changes: 7 additions & 0 deletions ptodsl/tests/support/docs_fragment_fixtures.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down
9 changes: 9 additions & 0 deletions ptodsl/tests/test_jit_compile.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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 <PIPE_ALL>" 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")
Expand Down
24 changes: 24 additions & 0 deletions test/lit/vpto/trap_vpto_llvm.pto
Original file line number Diff line number Diff line change
@@ -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<vector>} {
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()
Loading