triton.language.sum
- triton.language.sum = <function sum>
Returns the sum of all elements in the
inputtensor along the providedaxis- 参数:
input (Tensor) -- the input values
axis (int) -- the dimension along which the reduction should be done. If None, reduce all dimensions
keep_dims (bool) -- if true, keep the reduced dimensions with length 1
dtype (tl.dtype) -- the desired data type of the returned tensor. If specified, the input tensor is casted to
dtypebefore the operation is performed. This is useful for preventing data overflows. If not specified, integer and bool dtypes are upcasted totl.int32and float dtypes are upcasted to at leasttl.float32.
This function can also be called as a member function on
tensor, asx.sum(...)instead ofsum(x, ...). .. rubric:: Exampleimport torch import torch_npu import triton import triton.language as tl @triton.jit def tt_sum_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.sum(x, dim) if dim == 0: tl.store(out_ptr + tl.arange(0, N), ret) else: tl.store(out_ptr + tl.arange(0, M), ret) def test_sum(): M, N, dim = 4, 8, 1 x = torch.randn(M, N, dtype=torch.float32).npu() out = torch.empty(M, dtype=torch.float32).npu() tt_sum_2d[1, 1, 1](x, out, M, N, dim) ref = torch.sum(x, dim=dim) assert torch.allclose(out.cpu(), ref.cpu()), "sum result mismatch" print("sum result:", out) if __name__ == "__main__": test_sum()
Special Restrictions
DataType: Ascend A2/A3 does not support uint16, uint32, uint64.
keep_dims=Truerequires more test coverage; currently verified for 3D tensor with dim=2.