triton.language.expand_dims

triton.language.expand_dims(input, axis, _semantic=None)

通过插入新的长度为 1 的维度来扩展张量的形状。

轴索引是相对于结果张量的,因此 result.shape[axis] 每个轴的值将为 1。

参数:
  • input (tl.tensor) – 输入张量。

  • axis (int | Sequence[int]) – 添加新轴的索引。

此函数也可以作为成员函数在 tensor 上调用,形式为 x.expand_dims(...) 而非 expand_dims(x, ...)