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
√
√
√
√
√
√
√
√
√
√
×
√
√
√
√