triton.language

编程模型

tensor

表示一个 N 维的值或指针数组。

tensor_descriptor

表示全局内存中张量的描述符。

program_id

返回当前程序实例在给定 axis(轴)上的 ID。

num_programs

返回在给定 axis(轴)上启动的程序实例数量。

创建操作

arange

返回半开区间 [start, end) 内的连续值。

cat

连接给定的块

full

返回一个由标量值填充的张量,具有给定的 shape(形状)和 dtype(数据类型)。

zeros

返回一个由标量值 0 填充的张量,具有给定的 shape(形状)和 dtype(数据类型)。

zeros_like

返回一个与给定张量形状和类型相同的全零张量。

cast

将张量转换为给定的 dtype

形状操作

broadcast

尝试将两个给定的块广播为兼容的公共形状。

broadcast_to

尝试将给定张量广播为新的 shape

expand_dims

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

interleave

沿最后一个维度交错两个张量的值。

join

在一个新的次要维度上连接给定的张量。

permute

置换张量的维度。

ravel

返回 x 的连续扁平视图。

reshape

返回一个与输入元素数量相同但形状不同的张量。

split

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

trans

置换张量的维度。

view

返回一个与 input 元素相同但形状不同的张量。

线性代数操作

dot

返回两个块的矩阵乘积。

dot_scaled

返回微缩放(microscaling)格式下两个块的矩阵乘积。

内存/指针操作

load

返回一个张量,其值是从 pointer 定义的内存位置加载的。

store

将一个张量存入 pointer 定义的内存位置。

make_tensor_descriptor

创建一个张量描述符对象。

load_tensor_descriptor

从张量描述符加载一个数据块。

store_tensor_descriptor

将一个数据块存入张量描述符。

make_block_ptr

返回指向父张量中某个块的指针。

advance

推进块指针。

索引操作

flip

沿维度 dim 翻转张量 x

where

根据 condition 返回从 xy 中选取的元素组成的张量。

swizzle2d

将行主序 size_i * size_j 矩阵的索引转换为每组 size_g 行的列主序矩阵索引。

数学操作

abs

计算 x 的逐元素绝对值。

cdiv

计算 x 除以 div 的向上取整除法。

ceil

计算 x 的逐元素向上取整值。

clamp

将输入张量 x 的值限制在 [min, max] 范围内。

cos

计算 x 的逐元素余弦值。

div_rn

计算 xy 的逐元素精确除法(根据 IEEE 标准舍入到最近值)。

erf

计算 x 的逐元素误差函数。

exp

计算 x 的逐元素指数。

exp2

计算 x 的逐元素指数(以 2 为底)。

fdiv

计算 xy 的逐元素快速除法。

floor

计算 x 的逐元素向下取整值。

fma

计算 xyz 的逐元素融合乘加(fused multiply-add)。

log

计算 x 的逐元素自然对数。

log2

计算 x 的逐元素以 2 为底的对数。

maximum

计算 xy 的逐元素最大值。

minimum

计算 xy 的逐元素最小值。

rsqrt

计算 x 的逐元素平方根倒数。

sigmoid

计算 x 的元素级 Sigmoid。

sin

计算 x 的逐元素正弦值。

softmax

计算 x 的元素级 Softmax。

sqrt

计算 x 的逐元素快速平方根。

sqrt_rn

计算 x 的逐元素精确平方根(根据 IEEE 标准舍入到最近值)。

umulhi

计算 xy 的 2N 位乘积中最重要的 N 位。

归约操作

argmax

返回 input 张量中沿给定 axis 的所有元素的最大索引。

argmin

返回 input 张量中沿给定 axis 的所有元素的最小索引。

最大值 (max)

返回 input 张量中沿给定 axis 的所有元素的最大值。

最小值 (min)

返回 input 张量中沿给定 axis 的所有元素的最小值。

reduce

沿提供的 axis(轴)将 combine_fn 应用于 input 张量中的所有元素。

sum

返回 input 张量中沿给定 axis 的所有元素的和。

xor_sum

返回 input 张量中沿给定 axis 的所有元素的异或和。

扫描/排序操作

associative_scan

沿提供的 axis(轴)将 combine_fn 应用于每个带有进位(carry)的 input 张量元素,并更新进位。

cumprod

返回 input 张量中沿给定 axis 的所有元素的累积积。

cumsum

返回 input 张量中沿给定 axis 的所有元素的累积和。

histogram

基于带有 num_bins 个桶的输入张量计算直方图,桶宽度为 1,从 0 开始。

sort

topk

返回输入张量在指定维度上最大的(或最小的)k 个元素。

gather

沿给定维度从张量中收集(gather)数据。

原子操作

atomic_add

pointer 指定的内存位置执行原子加法。

atomic_and

pointer 指定的内存位置执行原子逻辑与(atomic logical and)操作。

atomic_cas

pointer 指定的内存位置执行原子比较并交换(compare-and-swap)。

atomic_max

pointer 指定的内存位置执行原子取最大值。

atomic_min

pointer 指定的内存位置执行原子取最小值。

atomic_or

pointer 指定的内存位置执行原子逻辑或(logical or)。

atomic_xchg

pointer 指定的内存位置执行原子交换操作。

atomic_xor

pointer 指定的内存位置执行原子逻辑异或(logical xor)。

随机数生成

randint4x

给定 seed(标量)和 offset(块),返回四个随机 int32 块。

randint

给定 seed(标量)和 offset(块),返回一个随机 int32 块。

rand

给定 seed(标量)和 offset(块),返回一个 \(U(0, 1)\) 分布的随机 float32 块。

randn

给定 seed(标量)和 offset(块),返回一个 \(\mathcal{N}(0, 1)\) 分布的随机 float32 块。

迭代器

range

永远向上计数的迭代器。

static_range

永远向上计数的迭代器。

内联汇编

inline_asm_elementwise

在张量上执行内联汇编。

编译器提示操作

assume

允许编译器假定 cond 为 True。

debug_barrier

插入栅栏(barrier)以同步块中的所有线程。

max_constancy

让编译器知道 input 中前 value 个值是恒定的。

max_contiguous

让编译器知道 input 中前 value 个值是连续的。

multiple_of

让编译器知道 input 中的值都是 value 的倍数。

调试操作

static_print

在编译时打印值。

static_assert

在编译时断言条件。

device_print

在运行时从设备打印值。

device_assert

在运行时从设备断言条件。