Triton 语义
Triton 在很大程度上遵循 NumPy 的语义,仅有少数例外。在本文档中,我们将介绍 Triton 支持的一些数组计算特性,并涵盖 Triton 语义与 NumPy 存在偏差的情况。
类型提升
类型提升 (Type Promotion) 发生在运算中使用不同数据类型的张量时。对于与双下划线方法 (dunder methods) 关联的二元运算,以及三元函数 tl.where 的后两个参数,Triton 会根据种类层次结构(数据类型集合)自动将输入张量转换为通用数据类型:{bool} < {integral dypes} < {floating point dtypes}。
算法规则如下
种类 (Kind) 如果一个张量的数据类型属于更高阶的种类,则另一个张量会被提升至该类型:
(int32, bfloat16) -> bfloat16宽度 (Width) 如果两个张量的数据类型属于同一种类,但其中一个具有更高的位宽,则另一个会被提升至该类型:
(float32, float16) -> float32优先选择 float16 如果两个张量具有相同的宽度和符号性质,但数据类型不同(如
float16和bfloat16,或者不同的fp8类型),它们都会被提升至float16。(float16, bfloat16) -> float16优先选择无符号 其他情况(宽度相同,符号性质不同),它们会被提升至无符号数据类型:
(int32, uint32) -> uint32
当涉及标量时,规则有所不同。此处的“标量”指数字字面量、被标记为 tl.constexpr 的变量,或这些项的组合。它们由 NumPy 标量表示,类型包括 bool、int 和 float。
当运算涉及张量和标量时
如果标量的种类低于或等于张量,它将不参与提升:
(uint8, int) -> uint8如果标量的种类更高,我们将选择它能适配的最低数据类型,整数范围为
int32<uint32<int64<uint64,浮点数范围为float32<float64。然后,张量和标量都会被提升至该类型:(int16, 4.0) -> float32
广播机制
广播 (Broadcasting) 允许对不同形状的张量进行运算,通过自动将它们的形状扩展为兼容大小而无需复制数据。其遵循以下规则:
如果其中一个张量的形状较短,则在其左侧填充 1,直到两个张量的维度数量相同:
((3, 4), (5, 3, 4)) -> ((1, 3, 4), (5, 3, 4))如果两个维度相等,或者其中一个为 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 语义。