triton.language.reduce

triton.language.reduce(input, axis, combine_fn, keep_dims=False, _semantic=None, _generator=None)

Applies the combine_fn to all elements in input tensors along the provided axis

参数:
  • input (Tensor) -- the input tensor, or tuple of tensors

  • axis (int | None) -- the dimension along which the reduction should be done. If None, reduce all dimensions

  • combine_fn (Callable) -- a function to combine two groups of scalar tensors (must be marked with @triton.jit)

  • keep_dims -- if true, keep the reduced dimensions with length 1

Example

import torch
import torch_npu
import triton
import triton.language as tl


@triton.jit
def _reduce_combine(a, b):
    return a + b


@triton.jit
def tt_reduce_2d(in_ptr, out_ptr, M: tl.constexpr, N: tl.constexpr, dim: tl.constexpr):
    mblk_idx = tl.arange(0, M)
    nblk_idx = tl.arange(0, N)
    idx = mblk_idx[:, None] * N + nblk_idx[None, :]
    x = tl.load(in_ptr + idx)
    ret = tl.reduce(x, dim, _reduce_combine)
    if dim == 0:
        tl.store(out_ptr + tl.arange(0, N), ret)
    else:
        tl.store(out_ptr + tl.arange(0, M), ret)


def test_reduce():
    M, N, dim = 4, 8, 1
    x = torch.randn(M, N, dtype=torch.float32).npu()
    out = torch.empty(M, dtype=torch.float32).npu()
    tt_reduce_2d[1, 1, 1](x, out, M, N, dim)
    ref = torch.sum(x, dim=dim)
    assert torch.allclose(out.cpu(), ref.cpu()), "reduce result mismatch"
    print("reduce result:", out)


if __name__ == "__main__":
    test_reduce()

Special Restrictions

  • DataType: Ascend A2/A3 does not support uint16, uint32, uint64; Ascend 950 does not fp8e4(E4M3), fp8e5(E5M2) (hardware limitation).

  • keep_dims=True requires more test coverage; currently verified for 3D tensor with dim=2.