triton.language.split

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

沿张量的最后一个维度将张量拆分为两部分,最后一个维度的大小必须为 2。

例如,给定形状为 (4,8,2) 的张量,生成两个形状为 (4,8) 的张量。给定形状为 (2) 的张量,返回两个标量。

如果你想拆分成多于两部分,你可以多次调用此函数(可能还需要调用 reshape)。这反映了 Triton 中张量大小必须是 2 的幂的约束。

split 是 join 的逆操作。

参数:

a (Tensor) – 要拆分的张量。

此函数也可以作为成员函数在 tensor 上调用,例如 x.split() 而不是 split(x)