triton.language.join

Contents

triton.language.join#

triton.language.join(a, b, _semantic=None)#

Join the given tensors in a new, minor dimension.

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

The two inputs are broadcasted to be the same shape.

If you want to join more than two elements, you can use multiple calls to this function. This reflects the constraint in Triton that tensors must have power-of-two sizes.

join is the inverse of split.

Parameters:
  • a (Tensor) – The first input tensor.

  • b (Tensor) – The second input tensor.

Example

import torch
import triton
import triton.language as tl


@triton.jit
def join_2d_kernel(x_ptr, y_ptr, out_ptr, M: tl.constexpr, N: tl.constexpr):
    # (M, N) -> (2, M, N)
    offsets_m = tl.arange(0, M)[:, None]
    offsets_n = tl.arange(0, N)[None, :]
    x = tl.load(x_ptr + offsets_m * N + offsets_n)
    y = tl.load(y_ptr + offsets_m * N + offsets_n)
    z = tl.join(x, y)
    flat = tl.reshape(z, (2 * M * N, ))
    tl.store(out_ptr + tl.arange(0, 2 * M * N), flat)


def test_join():
    M, N = 2, 3
    x = torch.zeros([M, N], dtype=torch.float32).npu()
    y = torch.zeros([M, N], dtype=torch.float32).npu()

    out = torch.empty((2 * M * N, ), dtype=torch.float32, device="npu")
    join_2d_kernel[(1, )](x, y, out, M=M, N=N)

    assert out.shape == (2 * M * N, ), f"Shape mismatch: expected {(2*M*N,)} but got {out.shape}."
    print("Test passed.")


if __name__ == "__main__":
    test_join()

DataType Support

平台

uint8

int8

uint16

int16

uint32

int32

uint64

int64

fp16

fp32

fp64

bf16

fp8e(e4m3)

fp8e5(e5m2)

bool

Ascend A2/A3

×

×

×

×

×

×

Ascend 950

×