Triton 语义

Triton 在很大程度上遵循 NumPy 的语义,仅有少数例外。在本文档中,我们将介绍 Triton 支持的一些数组计算特性,并涵盖 Triton 语义与 NumPy 存在偏差的情况。

类型提升

类型提升 (Type Promotion) 发生在运算中使用不同数据类型的张量时。对于与双下划线方法 (dunder methods) 关联的二元运算,以及三元函数 tl.where 的后两个参数,Triton 会根据种类层次结构(数据类型集合)自动将输入张量转换为通用数据类型:{bool} < {integral dypes} < {floating point dtypes}

算法规则如下

  1. 种类 (Kind) 如果一个张量的数据类型属于更高阶的种类,则另一个张量会被提升至该类型:(int32, bfloat16) -> bfloat16

  2. 宽度 (Width) 如果两个张量的数据类型属于同一种类,但其中一个具有更高的位宽,则另一个会被提升至该类型:(float32, float16) -> float32

  3. 优先选择 float16 如果两个张量具有相同的宽度和符号性质,但数据类型不同(如 float16bfloat16,或者不同的 fp8 类型),它们都会被提升至 float16(float16, bfloat16) -> float16

  4. 优先选择无符号 其他情况(宽度相同,符号性质不同),它们会被提升至无符号数据类型:(int32, uint32) -> uint32

当涉及标量时,规则有所不同。此处的“标量”指数字字面量、被标记为 tl.constexpr 的变量,或这些项的组合。它们由 NumPy 标量表示,类型包括 boolintfloat

当运算涉及张量和标量时

  1. 如果标量的种类低于或等于张量,它将不参与提升:(uint8, int) -> uint8

  2. 如果标量的种类更高,我们将选择它能适配的最低数据类型,整数范围为 int32 < uint32 < int64 < uint64,浮点数范围为 float32 < float64。然后,张量和标量都会被提升至该类型:(int16, 4.0) -> float32

广播机制

广播 (Broadcasting) 允许对不同形状的张量进行运算,通过自动将它们的形状扩展为兼容大小而无需复制数据。其遵循以下规则:

  1. 如果其中一个张量的形状较短,则在其左侧填充 1,直到两个张量的维度数量相同:((3, 4), (5, 3, 4)) -> ((1, 3, 4), (5, 3, 4))

  2. 如果两个维度相等,或者其中一个为 1,则它们是兼容的。维度 1 将被扩展以匹配另一个张量的维度。((1, 3, 4), (5, 3, 4)) -> ((5, 3, 4), (5, 3, 4))

与 NumPy 的差异

整数除法中的 C 语言舍入规则 为了提高效率,Triton 中的运算符遵循 C 语言语义而非 Python 语义。因此,int // int 在处理不同符号的整数时,实现的是向零舍入(类似 C 语言),而不是 Python 中的向负无穷舍入。出于同样的原因,取模运算符 int % int(定义为 a % b = a - b * (a // b))也遵循 C 语言语义而非 Python 语义。

可能令人困惑的是,对于所有输入均为标量的计算,整数除法和取模遵循 Python 语义。