From 5bfedde4664351e5b6a672aa26d460c4479eabec Mon Sep 17 00:00:00 2001 From: ggbondbest <200944273+ggbondbest@users.noreply.github.com> Date: Fri, 2 Oct 2026 21:56:00 +0800 Subject: [PATCH 1/2] Add minimal Ascend INT8 ND2NZ TLE primitive --- .../dsa/ascend/custom_ops/CUSTOM_OP_USAGE.md | 60 ++++++++ .../dsa/ascend/custom_ops/__init__.py | 2 + .../dsa/ascend/custom_ops/build_custom_ops.sh | 1 + .../mem_ops/data_copy_gm_to_l1_nd2nz_int8.cpp | 19 +++ .../dsa/ascend/custom_ops/registry.py | 54 +++++++ .../tutorials/tle/custom/test_custom_ops.py | 3 + .../tutorials/tle/custom/test_nd2nz_int8.py | 143 ++++++++++++++++++ .../tle/dsa/dialect/lib/CMakeLists.txt | 1 + 8 files changed, 283 insertions(+) create mode 100644 python/triton/experimental/tle/language/dsa/ascend/custom_ops/mem_ops/data_copy_gm_to_l1_nd2nz_int8.cpp create mode 100644 python/tutorials/tle/custom/test_nd2nz_int8.py diff --git a/python/triton/experimental/tle/language/dsa/ascend/custom_ops/CUSTOM_OP_USAGE.md b/python/triton/experimental/tle/language/dsa/ascend/custom_ops/CUSTOM_OP_USAGE.md index 6f9a393020..f0167fc60e 100644 --- a/python/triton/experimental/tle/language/dsa/ascend/custom_ops/CUSTOM_OP_USAGE.md +++ b/python/triton/experimental/tle/language/dsa/ascend/custom_ops/CUSTOM_OP_USAGE.md @@ -11,6 +11,7 @@ custom_ops/ ├── build_custom_ops.sh # 手动重新编译统一 bitcode ├── custom_ops.bc # 所有注册算子共用的 bitcode,编译后生成 ├── mem_ops/ +│ ├── data_copy_gm_to_l1_nd2nz_int8.cpp # Single INT8 GM → L1 ND2NZ transfer │ ├── duplicate.cpp # 将一个变量或立即数复制多次并填充到向量中,暂只支持 tensor 高维切分计算中 mask 逐比特模式 │ ├── gather_gm_to_l1.cpp # GM → L1/CBUF 按索引行 gather │ ├── gather_gm_to_ub.cpp # GM → UB 按索引行 gather @@ -66,6 +67,7 @@ output0, output1 = tle.dsa.ascend.raw( | --- | --- | --- | --- | --- | | `duplicate_bitwise_mask` | VECTOR / V | 将一个变量或立即数复制多次并填充到向量中 | 需要填充数据的向量,对应 Ascend C `const LocalTensor& dst` | `mem_ops/ duplicate.cpp` | | `gather_gm_to_l1` | CUBE / MTE2 | 按索引将 GM 连续张量中的 half/bf16 数据行收集到 L1/CBUF,并完成 ND2NZ 搬运 | L1/CBUF half/bf16 目标张量,对应 C++ `dst` | `mem_ops/gather_gm_to_l1.cpp` | +| `data_copy_gm_to_l1_nd2nz_int8` | CUBE / MTE2 | One signed INT8 GM → L1 ND2NZ transfer | Logical 2D INT8 output, consumed by `tl.dot` as 4D L1 NZ | `mem_ops/data_copy_gm_to_l1_nd2nz_int8.cpp` | | `gather_gm_to_ub` | VECTOR / MTE2 | 按索引将 GM 连续张量中的 half/bf16 数据行收集到 UB | UB half/bf16 目标张量,对应 C++ `dst` | `mem_ops/gather_gm_to_ub.cpp` | | `gather_mask_builtin_pattern` | VECTOR / V | 以内置固定模式对应的二进制对应的二进制为 gather mask, 从源操作数中选取元素写入目的操作数中 | 目的操作数,`out[0]` 对应 Ascend C `const LocalTensor& dst`, `out[1]` 对应 Ascend C `uint64_t& rsvdCnt` | `mem_ops/gather_mask.cpp` | | `gather_mask_custom_pattern` | VECTOR / V | 以用户自定义输入的 Tensor 数值对应的二进制为 gather mask, 从源操作数中选取元素写入目的操作数中 | 目的操作数,`out[0]`对应 Ascend C `const LocalTensor& dst`, `out[1]`对应 Ascend C `uint64_t& rsvdCnt` | `mem_ops/gather_mask.cpp` | @@ -101,6 +103,64 @@ dst = tle.dsa.ascend.raw( 完整示例见 `python/tutorials/tle/custom/test_custom_ops.py`(`test_duplicate`)。 +### `data_copy_gm_to_l1_nd2nz_int8` + +```python +tile = tle.dsa.ascend.raw( + "data_copy_gm_to_l1_nd2nz_int8", + src, nd_num, n_value, d_value, src_nd_matrix_stride, src_d_value, + dst_nz_c0_stride, dst_nz_n_stride, dst_nz_matrix_stride, + out=tile, +) +``` + +This is the signed INT8 overload of CANN's +`DataCopy(dstLocal, srcGlobal, Nd2NzParams)`. All eight dav_c220 struct fields +are exposed with their original `uint16_t` types. Static parameters are +checked against CANN 9.1 `CheckNd2NzParamsCommon`: + +| Parameter | Range | Meaning | +| --- | --- | --- | +| `nd_num` | 0–4095 | Number of source matrices | +| `n_value` | 0–16384 | Rows per source matrix | +| `d_value` | 0–65535 | Columns per source matrix | +| `src_nd_matrix_stride` | 0–65535 | Source matrix stride, in INT8 elements | +| `src_d_value` | 1–65535 | Source row stride, in INT8 elements | +| `dst_nz_c0_stride` | 1–16384 | Destination stride between C0 blocks, in 32-byte units | +| `dst_nz_n_stride` | 1–16384 | Destination row stride, in 32-byte units | +| `dst_nz_matrix_stride` | 0–65535 | Destination matrix stride, in INT8 elements | + +`src` must be a rank-two GM block pointer and `out` a logical rank-two INT8 +tensor. A following `tl.dot` anchors its physical L1 NZ layout to +`[ceil(cols / 32), ceil(rows / 16), 16, 32]`; a standalone raw-output-to-store +graph is not supported. Static destination strides are checked against that +buffer's capacity using CANN's overflow formula. Dynamic parameters must +satisfy the same ranges and bounds at runtime. Source bounds, source strides, +destination alignment, and initialization of any untouched elements remain +the caller's responsibility. Zero `nd_num`, `n_value`, or `d_value` is a no-op. + +For an aligned `[M, K]` tile with contiguous columns and a valid `row_stride`, +the parameter tuple is `(1, M, K, 0, row_stride, M, 1, 1)`. The source block +pointer carries the starting offset. A row stride above 65535 must be handled +outside this primitive, rather than narrowed to uint16. + +The C++ implementation follows CANN 9.1 +`dav_c220/kernel_operator_data_copy_impl.h::DataCopyGM2L1ND2NZImplBase` and +calls `copy_gm_to_cbuf_multi_nd2nz_b8` once. It contains no allocation, loop, +matrix multiply, scale application, or pipeline barrier. Pipeline ordering +is the caller/compiler's responsibility. On CANN 9.1, use +`enable_legacy_insert_load_store_for_mix_cv=True` for the same custom MTE2 +output-memory-scope inference issue documented for `gather_gm_to_l1` below. + +`python/tutorials/tle/custom/test_nd2nz_int8.py` checks registration and exact +INT8 copy results through an identity matrix multiply, including source +padding, nonzero offsets, and output guards. Device cases use `nd_num=1`. +Dynamic parameters have registration-only coverage; multiple-matrix copies +have not been device-validated. Its optional `--benchmark` +compares this primitive with `dsa.alloc/copy/to_tensor` using identical +identity-dot consumers and compiler options. It does not benchmark a whole +quantized operator or hide application packing costs. + ### `gather_gm_to_l1` ```python diff --git a/python/triton/experimental/tle/language/dsa/ascend/custom_ops/__init__.py b/python/triton/experimental/tle/language/dsa/ascend/custom_ops/__init__.py index 9c7dd1dccc..9c8414eda2 100644 --- a/python/triton/experimental/tle/language/dsa/ascend/custom_ops/__init__.py +++ b/python/triton/experimental/tle/language/dsa/ascend/custom_ops/__init__.py @@ -9,6 +9,7 @@ # 面向用户的 custom op from .registry import ( duplicate_bitwise_mask, + data_copy_gm_to_l1_nd2nz_int8, gather_gm_to_l1, gather_gm_to_ub, gather_mask_builtin_pattern, @@ -28,6 +29,7 @@ "SORT_IMPL_S4096_K129_512", "SORT_IMPL_S4096_K1_128_K2048", "duplicate_bitwise_mask", + "data_copy_gm_to_l1_nd2nz_int8", "gather_gm_to_l1", "gather_gm_to_ub", "gather_mask_builtin_pattern", diff --git a/python/triton/experimental/tle/language/dsa/ascend/custom_ops/build_custom_ops.sh b/python/triton/experimental/tle/language/dsa/ascend/custom_ops/build_custom_ops.sh index 2a90a9201e..82b0a2ff6b 100755 --- a/python/triton/experimental/tle/language/dsa/ascend/custom_ops/build_custom_ops.sh +++ b/python/triton/experimental/tle/language/dsa/ascend/custom_ops/build_custom_ops.sh @@ -34,6 +34,7 @@ fi # bitcode: "::" — arch is the ccec aicore target for that op. CUSTOM_OPS=( "mem_ops/duplicate.cpp:dav-c220-vec" + "mem_ops/data_copy_gm_to_l1_nd2nz_int8.cpp:dav-c220-cube" "mem_ops/gather_gm_to_l1.cpp:dav-c220-cube" "mem_ops/gather_gm_to_ub.cpp:dav-c220-vec" "mem_ops/gather_mask.cpp:dav-c220-vec" diff --git a/python/triton/experimental/tle/language/dsa/ascend/custom_ops/mem_ops/data_copy_gm_to_l1_nd2nz_int8.cpp b/python/triton/experimental/tle/language/dsa/ascend/custom_ops/mem_ops/data_copy_gm_to_l1_nd2nz_int8.cpp new file mode 100644 index 0000000000..7bb8c24a53 --- /dev/null +++ b/python/triton/experimental/tle/language/dsa/ascend/custom_ops/mem_ops/data_copy_gm_to_l1_nd2nz_int8.cpp @@ -0,0 +1,19 @@ +// Copyright 2026 FlagOS Contributors +// SPDX-License-Identifier: Apache-2.0 +#include "Utils.h" + +// Follows CANN 9.1 dav_c220/kernel_operator_data_copy_impl.h, +// DataCopyGM2L1ND2NZImplBase, int8_t branch. All dav_c220 Nd2NzParams +// fields are exposed. This primitive performs exactly one transfer; the +// caller owns bounds, padding, layout, and pipeline synchronization. +extern "C" __aicore__ __attribute__((always_inline)) void +_mlir_ciface_custom_data_copy_gm_to_l1_nd2nz_int8( + memref_t<__gm__ int8_t, 2> *src, uint16_t nd_num, uint16_t n_value, + uint16_t d_value, uint16_t src_nd_matrix_stride, uint16_t src_d_value, + uint16_t dst_nz_c0_stride, uint16_t dst_nz_n_stride, + uint16_t dst_nz_matrix_stride, memref_t<__cbuf__ int8_t, 4> *dst) { + copy_gm_to_cbuf_multi_nd2nz_b8( + dst->aligned + dst->offset, src->aligned + src->offset, 0, nd_num, + n_value, d_value, src_nd_matrix_stride, src_d_value, dst_nz_c0_stride, + dst_nz_n_stride, dst_nz_matrix_stride); +} diff --git a/python/triton/experimental/tle/language/dsa/ascend/custom_ops/registry.py b/python/triton/experimental/tle/language/dsa/ascend/custom_ops/registry.py index f9253ea79c..3ea3a9e776 100644 --- a/python/triton/experimental/tle/language/dsa/ascend/custom_ops/registry.py +++ b/python/triton/experimental/tle/language/dsa/ascend/custom_ops/registry.py @@ -66,6 +66,60 @@ def __init__(self, scalar_value, mask, repeat_times, dst_block_stride, dst_repea self.extra_buffers = [(tl.float16, 0)] +@al.register_custom_op +class data_copy_gm_to_l1_nd2nz_int8: + """One signed INT8 GM-to-L1 DataCopy with all dav_c220 Nd2NzParams. + + The caller owns source bounds, ND/NZ layouts and pipeline ordering. + A following tl.dot anchors the logical rank-two output's rank-four L1 + NZ layout. A standalone raw-output-to-store graph is not supported. + Static parameters follow CANN CheckNd2NzParamsCommon; dynamic parameters + must satisfy those same bounds at runtime. + """ + + core = al.CORE.CUBE + pipe = al.PIPE.PIPE_MTE2 + mode = al.MODE.SIMD + + def __init__(self, src, nd_num: tl.uint16, n_value: tl.uint16, d_value: tl.uint16, src_nd_matrix_stride: tl.uint16, + src_d_value: tl.uint16, dst_nz_c0_stride: tl.uint16, dst_nz_n_stride: tl.uint16, + dst_nz_matrix_stride: tl.uint16, out=None): + assert out is not None, "data_copy_gm_to_l1_nd2nz_int8 requires an output buffer" + assert _element_dtype(src) == tl.int8, "src must contain signed INT8" + assert _element_dtype(out) == tl.int8, "out must contain signed INT8" + assert src.dtype.is_ptr() and src.dtype.element_ty.is_block(), ( + "src must be a two-dimensional GM block pointer") + assert len(src.dtype.element_ty.shape) == 2, "src must have rank two" + assert len(out.shape) == 2, "out must have logical rank two" + # CANN 9.1 kernel_operator_data_copy_check.h, CheckNd2NzParamsCommon. + fields = ( + ("nd_num", nd_num, 0, 4095), + ("n_value", n_value, 0, 16384), + ("d_value", d_value, 0, 65535), + ("src_nd_matrix_stride", src_nd_matrix_stride, 0, 65535), + ("src_d_value", src_d_value, 1, 65535), + ("dst_nz_c0_stride", dst_nz_c0_stride, 1, 16384), + ("dst_nz_n_stride", dst_nz_n_stride, 1, 16384), + ("dst_nz_matrix_stride", dst_nz_matrix_stride, 0, 65535), + ) + for name, value, lower, upper in fields: + assert not isinstance(value, bool), f"{name} must be an integer scalar, not bool" + if isinstance(value, int): + assert lower <= value <= upper, f"{name} must be in [{lower}, {upper}] on dav_c220" + else: + assert isinstance(value, tl.tensor) and not value.type.is_block() and value.dtype.is_int(), ( + f"{name} must be an integer scalar") + if all(isinstance(value, int) for _, value, _, _ in fields) and nd_num and n_value and d_value: + # Match CANN's CheckDataCopyTensorSizeOverflow for signed INT8. + dst_bytes = ((nd_num - 1) * dst_nz_matrix_stride + (n_value - 1) * dst_nz_n_stride * 32 + + ((d_value + 31) // 32 - 1) * dst_nz_c0_stride * 32 + 32) + rows, cols = (int(dim) for dim in out.shape) + nz_bytes = ((rows + 15) // 16 * 16) * ((cols + 31) // 32 * 32) + assert dst_bytes <= nz_bytes, "ND2NZ destination strides exceed the output NZ buffer" + self.symbol = "custom_data_copy_gm_to_l1_nd2nz_int8" + self.bitcode = CUSTOM_OPS_BITCODE + + @al.register_custom_op class gather_gm_to_l1: """ diff --git a/python/tutorials/tle/custom/test_custom_ops.py b/python/tutorials/tle/custom/test_custom_ops.py index f144fdb041..6d99221d89 100644 --- a/python/tutorials/tle/custom/test_custom_ops.py +++ b/python/tutorials/tle/custom/test_custom_ops.py @@ -7,6 +7,7 @@ custom_ops.bc, one group per op: - gather_gm_to_l1 (fp16 / bf16): verified through a following tl.dot + - data_copy_gm_to_l1_nd2nz_int8: checked with an INT8 identity dot - gather_gm_to_ub (fp16 / bf16): verified by storing the result to GM - sort_1d_pack (all three sort paths BASE / S4096_K129_512 / S4096_K1_128_K2048, plus an index_offset case) @@ -680,6 +681,7 @@ def test_duplicate(): def main(): from test_cast_ops import main as test_cast_ops from test_compare_scalar import main as test_compare_scalar + from test_nd2nz_int8 import main as test_nd2nz_int8 for torch_dtype, tol in ((torch.float16, 1e-3), (torch.bfloat16, 1e-2)): test_gather_gm_to_l1(torch_dtype, tol) @@ -689,6 +691,7 @@ def main(): test_unpack_sort() test_compare_scalar() test_cast_ops() + test_nd2nz_int8() test_sort32() test_mrgsort() test_gather_mask() diff --git a/python/tutorials/tle/custom/test_nd2nz_int8.py b/python/tutorials/tle/custom/test_nd2nz_int8.py new file mode 100644 index 0000000000..b5d3a11356 --- /dev/null +++ b/python/tutorials/tle/custom/test_nd2nz_int8.py @@ -0,0 +1,143 @@ +# Copyright 2026 FlagOS Contributors +# SPDX-License-Identifier: Apache-2.0 +"""Native INT8 ND2NZ correctness and an optional equivalent-TLE benchmark.""" + +import argparse + +import numpy as np +import torch +import torch_npu # noqa: F401 +import triton +import triton.experimental.tle as tle +import triton.language as tl +from triton.experimental.tle.language.dsa.ascend.custom_ops import data_copy_gm_to_l1_nd2nz_int8 + + +@triton.jit +def copy_dot_kernel(X, Identity, Out, ROW_STRIDE: tl.constexpr, SOURCE_OFFSET: tl.constexpr, M: tl.constexpr, + K: tl.constexpr, CUSTOM: tl.constexpr): + rows = tl.arange(0, M) + cols = tl.arange(0, K) + src = tl.make_block_ptr(X + SOURCE_OFFSET, (M, K), (ROW_STRIDE, 1), (0, 0), (M, K), (1, 0)) + if CUSTOM: + tile = tle.dsa.ascend.raw("data_copy_gm_to_l1_nd2nz_int8", src, 1, M, K, 0, ROW_STRIDE, M, 1, 1, out=tl.full( + (M, K), 0, tl.int8)) + else: + buf = tle.dsa.alloc([M, K], tl.int8, tle.dsa.ascend.L1) + tle.dsa.copy(src, buf, [M, K]) + tile = tle.dsa.to_tensor(buf) + identity = tl.load(Identity + cols[:, None] * K + cols[None, :]) + result = tl.dot(tile, identity, out_dtype=tl.int32) + tl.store(Out + rows[:, None] * K + cols[None, :], result) + + +def _inputs(m, k, row_stride, offset): + rng = np.random.default_rng(42) + values = rng.integers(-128, 128, size=offset + m * row_stride + 32, dtype=np.int8) + # Cover all encodings without making distant rows repeat modulo 256. + values[:256] = np.arange(-128, 128, dtype=np.int16).astype(np.int8) + expected = values[offset + np.arange(m)[:, None] * row_stride + np.arange(k)[None, :]].astype(np.int32) + x = torch.from_numpy(values).to("npu") + identity = torch.eye(k, dtype=torch.int32).to(torch.int8).to("npu") + storage = torch.full((m * k + 64, ), 777777, dtype=torch.int32, device="npu") + return x, identity, storage, expected + + +def test_data_copy_gm_to_l1_nd2nz_int8(): + cases = ((16, 32, 32, 0), (16, 128, 160, 32), (32, 128, 256, 64), (128, 128, 128, 32), (256, 128, 128, 64), + (16, 128, 65504, 32)) + for m, k, row_stride, offset in cases: + x, identity, storage, expected = _inputs(m, k, row_stride, offset) + copy_dot_kernel[(1, )](x, identity, storage[32:-32], row_stride, offset, m, k, True, + enable_legacy_insert_load_store_for_mix_cv=True) + torch.npu.synchronize() + actual = storage.cpu().numpy() + np.testing.assert_array_equal(actual[32:-32].reshape(m, k), expected) + np.testing.assert_array_equal(actual[:32], np.full(32, 777777)) + np.testing.assert_array_equal(actual[-32:], np.full(32, 777777)) + print(f"[PASS] data_copy_gm_to_l1_nd2nz_int8: {len(cases)} layout/offset cases") + + +def bench_data_copy_gm_to_l1_nd2nz_int8(): + m, k, row_stride, offset = 128, 128, 160, 32 + x, identity, storage, expected = _inputs(m, k, row_stride, offset) + + def launch(custom): + # Keep the compiler options and the identity-dot consumer identical. + return copy_dot_kernel[(1, )](x, identity, storage[32:-32], row_stride, offset, m, k, custom, + enable_legacy_insert_load_store_for_mix_cv=True) + + for custom in (False, True): + launch(custom) + torch.npu.synchronize() + np.testing.assert_array_equal(storage[32:-32].cpu().numpy().reshape(m, k), expected) + timings = {"tle": [], "custom": []} + for order in ((False, True), (True, False)): + for custom in order: + name = "custom" if custom else "tle" + timings[name].append(triton.testing.do_bench(lambda custom=custom: launch(custom), return_mode="median")) + tle_ms = sum(timings["tle"]) / len(timings["tle"]) + custom_ms = sum(timings["custom"]) / len(timings["custom"]) + print(f"[BENCH] INT8 ND2NZ + identity dot: TLE {tle_ms * 1000:.3f} us, " + f"custom {custom_ms * 1000:.3f} us, speedup {tle_ms / custom_ms:.3f}x; rounds_ms={timings}") + + +def _tensor(dtype, shape): + return tl.tensor(None, tl.block_type(dtype, shape)) + + +def _block_pointer(dtype, shape): + return tl.tensor(None, tl.pointer_type(tl.block_type(dtype, shape))) + + +def test_validation(): + src = _block_pointer(tl.int8, [128, 128]) + dst = _tensor(tl.int8, [128, 128]) + args = [src, 1, 128, 128, 0, 160, 128, 1, 1] + op = data_copy_gm_to_l1_nd2nz_int8(*args, out=dst) + assert op.symbol == "custom_data_copy_gm_to_l1_nd2nz_int8" + assert op.bitcode.endswith("custom_ops.bc") + # Dynamic scalar parameters retain the native API and its runtime bounds. + dynamic = [src] + [tl.tensor(None, tl.uint16) for _ in range(8)] + data_copy_gm_to_l1_nd2nz_int8(*dynamic, out=dst) + # Zero-size transfers are valid no-ops in the official API. + noop = [src, 0, 0, 0, 65535, 65535, 16384, 16384, 65535] + data_copy_gm_to_l1_nd2nz_int8(*noop, out=dst) + invalid = [] + for position, value in ((1, -1), (1, 4096), (2, 16385), (3, 65536), (4, 65536), (5, 0), (5, 114688), (6, 0), + (6, 16385), (7, 0), (7, 16385), (8, 65536), (6, 256), (7, 2), (2, 1.5)): + changed = args.copy() + changed[position] = value + invalid.append((changed, dst)) + # The custom-op argument converter lowers Python bool to i1 before applying + # its declared type, which cannot match this primitive's uint16 ABI. + for position in range(1, 9): + changed = args.copy() + changed[position] = True + invalid.append((changed, dst)) + for bad_src in (_block_pointer(tl.float16, [128, 128]), _block_pointer(tl.int8, [16384]), + tl.tensor(None, tl.pointer_type(tl.int8)), _tensor(tl.int8, [128, 128])): + invalid.append(([bad_src] + args[1:], dst)) + for bad_dst in (None, _tensor(tl.float16, [128, 128]), _tensor(tl.int8, [16384]), _tensor(tl.int8, [16, 128])): + invalid.append((args, bad_dst)) + for bad_args, bad_dst in invalid: + try: + data_copy_gm_to_l1_nd2nz_int8(*bad_args, out=bad_dst) + except AssertionError: + continue + raise AssertionError("data_copy_gm_to_l1_nd2nz_int8 accepted an invalid signature") + print(f"[PASS] ND2NZ registration: three valid and {len(invalid)} rejected signatures") + + +def main(): + test_validation() + test_data_copy_gm_to_l1_nd2nz_int8() + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument("--benchmark", action="store_true") + args = parser.parse_args() + main() + if args.benchmark: + bench_data_copy_gm_to_l1_nd2nz_int8() diff --git a/third_party/tle/dsa/dialect/lib/CMakeLists.txt b/third_party/tle/dsa/dialect/lib/CMakeLists.txt index 78ed19fab1..af43e3aac7 100644 --- a/third_party/tle/dsa/dialect/lib/CMakeLists.txt +++ b/third_party/tle/dsa/dialect/lib/CMakeLists.txt @@ -31,6 +31,7 @@ set(ASCEND_CUSTOM_OPS_BC_DIR "${CMAKE_CURRENT_BINARY_DIR}/ascend_custom_ops") # bitcode: "::" — arch is the ccec aicore target for that op. set(ASCEND_CUSTOM_OPS "mem_ops/duplicate.cpp::dav-c220-vec" + "mem_ops/data_copy_gm_to_l1_nd2nz_int8.cpp::dav-c220-cube" "mem_ops/gather_gm_to_l1.cpp::dav-c220-cube" "mem_ops/gather_gm_to_ub.cpp::dav-c220-vec" "mem_ops/gather_mask.cpp::dav-c220-vec" From 1bed35fec16837737f78b7ba24aaa1a7e9522131 Mon Sep 17 00:00:00 2001 From: ggbondbest <200944273+ggbondbest@users.noreply.github.com> Date: Fri, 2 Oct 2026 22:51:48 +0800 Subject: [PATCH 2/2] Check static ND2NZ output bounds with dynamic source strides --- .../tle/language/dsa/ascend/custom_ops/registry.py | 3 ++- python/tutorials/tle/custom/test_nd2nz_int8.py | 6 ++++++ 2 files changed, 8 insertions(+), 1 deletion(-) diff --git a/python/triton/experimental/tle/language/dsa/ascend/custom_ops/registry.py b/python/triton/experimental/tle/language/dsa/ascend/custom_ops/registry.py index 3ea3a9e776..248eb88257 100644 --- a/python/triton/experimental/tle/language/dsa/ascend/custom_ops/registry.py +++ b/python/triton/experimental/tle/language/dsa/ascend/custom_ops/registry.py @@ -109,7 +109,8 @@ def __init__(self, src, nd_num: tl.uint16, n_value: tl.uint16, d_value: tl.uint1 else: assert isinstance(value, tl.tensor) and not value.type.is_block() and value.dtype.is_int(), ( f"{name} must be an integer scalar") - if all(isinstance(value, int) for _, value, _, _ in fields) and nd_num and n_value and d_value: + dst_fields = (nd_num, n_value, d_value, dst_nz_c0_stride, dst_nz_n_stride, dst_nz_matrix_stride) + if all(isinstance(value, int) for value in dst_fields) and nd_num and n_value and d_value: # Match CANN's CheckDataCopyTensorSizeOverflow for signed INT8. dst_bytes = ((nd_num - 1) * dst_nz_matrix_stride + (n_value - 1) * dst_nz_n_stride * 32 + ((d_value + 31) // 32 - 1) * dst_nz_c0_stride * 32 + 32) diff --git a/python/tutorials/tle/custom/test_nd2nz_int8.py b/python/tutorials/tle/custom/test_nd2nz_int8.py index b5d3a11356..b5363b1b1f 100644 --- a/python/tutorials/tle/custom/test_nd2nz_int8.py +++ b/python/tutorials/tle/custom/test_nd2nz_int8.py @@ -109,6 +109,12 @@ def test_validation(): changed = args.copy() changed[position] = value invalid.append((changed, dst)) + # A dynamic source stride must not bypass an entirely static output bound. + for position in (4, 5): + changed = args.copy() + changed[position] = tl.tensor(None, tl.uint16) + changed[6] = 256 + invalid.append((changed, dst)) # The custom-op argument converter lowers Python bool to i1 before applying # its declared type, which cannot match this primitive's uint16 ABI. for position in range(1, 9):