triton.language.split

Contents

triton.language.split#

triton.language.split(a, _semantic=None, _generator=None) tuple[tensor, tensor]#

Split a tensor in two along its last dim, which must have size 2.

For example, given a tensor of shape (4,8,2), produces two tensors of shape (4,8). Given a tensor of shape (2), returns two scalars.

If you want to split into more than two pieces, you can use multiple calls to this function (probably plus calling reshape). This reflects the constraint in Triton that tensors must have power-of-two sizes.

split is the inverse of join.

Parameters:

a – The tensor to split.

Example

import torch
import triton
import triton.language as tl


@triton.jit
def split_kernel(x_ptr, out0_ptr, out1_ptr, M: tl.constexpr):
    # (M, 2) -> (M, 1) + (M, 1)
    offsets_m = tl.arange(0, M)[:, None]
    offsets_n = tl.arange(0, 2)[None, :]
    x = tl.load(x_ptr + offsets_m * 2 + offsets_n)
    part0, part1 = x.split()
    flat0 = tl.reshape(part0, (M, ))
    flat1 = tl.reshape(part1, (M, ))
    tl.store(out0_ptr + tl.arange(0, M), flat0)
    tl.store(out1_ptr + tl.arange(0, M), flat1)


def test_split():
    M = 4
    x = torch.zeros([M, 2], dtype=torch.float32).npu()
    out0 = torch.empty((M, ), dtype=torch.float32, device="npu")
    out1 = torch.empty((M, ), dtype=torch.float32, device="npu")
    split_kernel[(1, )](out0, out1, x, M=M)

    assert out0.shape == (M, ) and out1.shape == (M, ), "Shape mismatch"
    print("Test passed.")


if __name__ == "__main__":
    test_split()

DataType Support

平台

uint8

int8

uint16

int16

uint32

int32

uint64

int64

fp16

fp32

fp64

bf16

fp8e(e4m3)

fp8e5(e5m2)

bool

Ascend A2/A3

×

×

×

×

×

×

Ascend 950

×