tinygrad 逐元素(Elementwise)运算全解:一元数学、激活函数、广播二元运算与类型转换
【免费下载链接】tinygradYou like pytorch? You like micrograd? You love tinygrad! ❤️项目地址: https://gitcode.com/GitHub_Trending/tiny/tinygrad
逐元素运算(Elementwise Ops)是 tinygrad 张量库中最基础、使用最频繁的一类操作:它们对张量的每个元素独立执行数学函数,不改变张量的形状。本文以 docs/tensor/elementwise.md 为骨架,结合 tinygrad/mixin/elementwise.py 与 tinygrad/uop/init.py 的源码实现,系统梳理 tinygrad 提供的全部一元数学运算、激活函数、广播二元/三元运算与类型转换操作,并深入讲解其背后的广播、类型提升与底层 UOp 机制,帮助你理解"每个算子背后发生了什么",从而在自定义算子或调试内核时游刃有余。
什么是 Elementwise 运算
Elementwise 运算的核心特征只有两条:
- 逐元素(per element):输出张量中每个位置的值只取决于输入张量中对应位置的值(以及可能的标量常量);
- 形状不变(shape-preserving):输入形状为
(d0, d1, ..., dn),输出形状仍然是(d0, d1, ..., dn)。
这一点在 tinygrad/uop/init.py 的GroupOp定义中有明确体现:Elementwise = ALU ∪ {CAST, BITCAST},其中ALU = Unary ∪ Binary ∪ Ternary。也就是说,tinygrad 在编译器内部把所有逐元素算子归为一类,统一走同一条代码生成与内核融合路径——这正是 tinygrad 能把y = a.relu().mul(b).add(c)这类表达式融合成单个 kernel 的基础。
底层机制:从 Tensor 方法到 UOp
文档列出的所有 API 最终都落在 tinygrad/mixin/elementwise.py 的ElementwiseMixin类中。理解三个核心辅助方法,就能理解全部算子:
1._broadcasted:广播与类型提升
def _broadcasted(self, y, reverse=False) -> tuple[Self, Self]: y = self.ufix(y) # 将 Python 常量包装成常量 UOp x, y = (self, y) if not reverse else (y, self) out_dtype = least_upper_dtype(x.dtype, y.dtype) # 计算"最小上界"目标类型 def promote(t): if t._uop.base.is_invalid: return t if t.dtype in dtypes.weaks and t._uop.base.op is Ops.CONST: return t if t.dtype == weak_dtype(out_dtype) else t._wrap_uop(remint(t._uop, dt)) return t.cast(out_dtype) return promote(x), promote(y)它做两件事:广播(把两个张量对齐到公共形状)和类型提升(通过least_upper_dtype计算出能同时容纳两个输入的最小精度类型,再各自cast过去)。例如int8 + float32的结果会被提升到float32;整数与浮点混合时遵循"弱类型"规则,保持常量不丢失精度。正是这个机制保证了文档中"Supports broadcasting to a common shape, type promotion"的承诺。
2._binop:二元运算统一入口
def _binop(self, op: Ops, x, reverse: bool) -> Self: lhs, rhs = self._broadcasted(x, reverse) return lhs.alu(op, rhs)所有二元运算(add、mul、maximum、bitwise_and等)都只是把底层Ops枚举值传给_binop而已。reverse参数用于支持右操作数形式(__radd__、__rmul__等)。
3.Ops枚举:运算的底层表示
tinygrad/uop/init.py 定义了与逐元素运算直接对应的底层算子:
- UnaryOps:
CAST、BITCAST、EXP2、LOG2、SIN、SQRT、RECIPROCAL、NEG、TRUNC - BinaryOps:
ADD、MUL、SHL、SHR、CDIV、MAX、CMOD、CMPLT、CMPNE、CMPEQ、XOR、OR、AND、THREEFRY、SUB、FDIV、POW、FLOORDIV、FLOORMOD - TernaryOps:
WHERE、MULACC
可以看到,很多高层 API(如sqrt、sin)是直接映射到硬件原语(Ops.SQRT、Ops.SIN),而另一些(如sigmoid、tanh)则是用更基础的原语组合出来的。下一节将逐个说明。
Unary Ops(数学一元运算)
文档将数学类一元运算与激活函数分开列出。数学类一元运算全部定义在ElementwiseMixin中,直接或间接映射到底层Ops:
| 方法 | 含义 | 底层实现要点 |
|---|---|---|
logical_not | 逻辑非(布尔取反) | cast(bool).ne(True) |
neg | 取负 | bool 走logical_not,否则self * (-1) |
log | 自然对数 | log2() * ln(2),即Ops.LOG2后乘常数 |
log2 | 以 2 为底的对数 | 直接Ops.LOG2 |
log10 | 以 10 为底的对数 | log2() * log10(2) |
exp | 自然指数 | cast(float32).mul(1/ln2).exp2(),最终落到Ops.EXP2 |
exp2 | 2 的幂 | 直接Ops.EXP2 |
sqrt | 平方根 | 直接Ops.SQRT |
rsqrt | 平方根倒数 | sqrt().reciprocal() |
sin/cos/tan | 三角函数 | sin直接Ops.SIN;cos用sin(π/2 - x)复合;tan为sin/cos |
asin/acos/atan | 反三角函数 | asin用多项式逼近(系数来自 Abramowitz & Stegun 4.4.46);acos = π/2 - asin;atan复合asin |
trunc | 向零截断 | 直接Ops.TRUNC |
ceil/floor | 向上/向下取整 | 基于trunc与where复合 |
round | 四舍五入(银行家舍入) | 实现保证 half-to-even 语义,与 NumPy 一致 |
isinf/isnan/isfinite | 数值状态检测 | isnan即self != self;isfinite为二者取反 |
lerp | 线性插值 | self + (end - self) * weight,uint8 有定点加速路径 |
square | 平方 | self * self |
clamp/clip | 数值截断 | clip是clamp的别名;支持单边None |
sign | 符号函数 | 基于where与常量 |
abs | 绝对值 | self * self.sign() |
reciprocal | 倒数 | 直接Ops.RECIPROCAL |
几个值得注意的实现细节:
cos不使用专门的Ops.COS。从源码可见cos通过sin(π/2 - x)计算,这是为了减少后端需要实现的硬件原语数量,让 CUDA/Metal/OpenCL 等各后端只需实现SIN一个三角函数即可覆盖cos、tan。round是"四舍六入五成双"(banker's rounding):-0.5 → -0.0、1.5 → 2.0,这一点与 Pythonround一致,但与某些语言向零舍入不同,使用时需注意。clamp/clip允许单边约束:min_或max_传None表示该侧无界,但不能同时为None(会抛RuntimeError)。
示例:
from tinygrad import Tensor x = Tensor([-3., -2., -1., 0., 1., 2., 3.]) print(x.abs().numpy()) # [3. 2. 1. 0. 1. 2. 3.] print(x.clamp(-1, 1).numpy()) # [-1. -1. -1. 0. 1. 1. 1.] print(Tensor([1., 2., 4., 8.]).log2().numpy()) # [0. 1. 2. 3.] print(Tensor([1, float('inf'), float('nan')]).isnan().numpy()) # [False False True]Unary Ops(激活函数)
激活函数类一元运算同样定义在ElementwiseMixin中,大部分由基础算子组合而成,没有引入新的硬件原语,因此可以在任意后端上运行:
| 方法 | 公式/含义 | 默认参数 | 实现要点 |
|---|---|---|---|
relu | max(x, 0) | — | (self > 0).where(self, 0) |
sigmoid | 1 / (1 + e^(-x)) | — | 基于exp2复合,避免引入EXP原语 |
logsigmoid | log(sigmoid(x)) | — | -(-x).softplus() |
hardsigmoid | 分段线性 sigmoid 近似 | alpha=1/6, beta=0.5 | 两个relu之差 |
elu | 指数线性单元 | alpha=1.0 | 分段where |
celu | 连续可微 ELU | alpha=1.0 | alpha * elu(x/alpha) |
selu | 缩放 ELU | alpha=1.67326, gamma=1.0507 | gamma * elu(alpha*x) |
swish/silu | x * sigmoid(x) | — | silu就是swish的别名 |
relu6 | min(max(x,0),6) | — | relu().minimum(6) |
hardswish | x * relu6(x+3) / 6 | — | 三个算子的组合 |
tanh | 双曲正切 | — | 2 * sigmoid(2x) - 1 |
sinh/cosh | 双曲正弦/余弦 | — | 基于exp组合 |
atanh/asinh/acosh | 反双曲函数 | — | 基于log、sqrt、square组合 |
hardtanh | 硬双曲正切 | min_val=-1, max_val=1 | 就是clip(min_val, max_val) |
erf | 误差函数 | — | Abramowitz & Stegun 7.1.26 多项式逼近 |
gelu | 高斯误差线性单元 | approximate="tanh" | 支持"tanh"与"none"两种模式 |
quick_gelu | x * sigmoid(1.702x) | — | Sigmoid GELU 近似 |
leaky_relu | 带泄漏的 ReLU | neg_slope=0.01 | (self < 0).where(neg_slope*self, self) |
mish | x * tanh(softplus(x)) | — | 组合实现 |
softplus | log(1 + e^x) | beta=1.0 | (1/beta) * logaddexp(beta*x, 0) |
softsign | x / (1 + |x|) | — | 组合实现 |
源码中有两个值得注意的"坑"被显式注释出来:
relu不能用self.maximum(0)实现:tinygrad/mixin/elementwise.py 注释说明maximum(0)在x == 0处会产生错误的梯度(会同时从两条路径回传一半梯度),因此relu刻意写成(self > 0).where(self, 0)以保证在 0 点处的梯度正确性。gelu默认使用 tanh 近似(approximate="tanh"),等价于 PyTorch 的默认approximate='tanh';传入"none"才使用基于erf的精确版本。
示例:
from tinygrad import Tensor x = Tensor([-3., -2., -1., 0., 1., 2., 3.]) print(x.relu().numpy()) # [0. 0. 0. 0. 1. 2. 3.] print(x.sigmoid().numpy()) print(x.gelu().numpy()) # tanh 近似 print(x.leaky_relu(neg_slope=0.42).numpy()) print(x.hardtanh(-0.5, 0.5).numpy()) # [-0.5 -0.5 -0.5 0. 0.5 0.5 0.5]Elementwise Ops(广播二元/三元运算)
这一类运算接受两个张量(或张量与标量),通过_broadcasted自动广播到公共形状并做类型提升。文档列出的全部 API:
| 方法 | 运算符 | 底层 Ops | 说明 |
|---|---|---|---|
add | + | ADD | 加法 |
sub | - | ADD(对取反后的 b) | 减法 |
mul | * | MUL | 乘法 |
div | / | FDIV复合 | 支持rounding_mode="trunc"/"floor" |
mod | % | FLOORMOD复合 | Python 风格 floor 取余 |
fmod | — | CMOD复合 | C 风格截断取余,符号跟随被除数 |
bitwise_xor | ^ | XOR | 按位异或 |
bitwise_and | & | AND | 按位与 |
bitwise_or | \| | OR | 按位或 |
bitwise_not | ~ | 复合 | 按位非(无符号取dtype.max异或,有符号异或-1) |
lshift/rshift | <</>> | SHL/SHR | 算术移位,要求整型 |
pow | ** | POW | 幂运算,支持reverse(如2.0 ** t) |
maximum | — | MAX | 逐元素最大值 |
minimum | — | MAX复合 | 逐元素最小值(有符号整型用 XOR 技巧实现) |
where | — | WHERE | 三元选择x_i if cond_i else y_i |
copysign | — | 复合 | 取self的幅值、other的符号 |
logaddexp | — | 复合 | 数值稳定的log(e^a + e^b) |
实现细节补充:
div的整数语义:默认对整数做"真除法"(先提升为浮点再除);传入rounding_mode="trunc"或"floor"时分别映射到底层CDIV/FLOORDIV,从而避免浮点路径。modvsfmod的区别:mod是 Python 风格的 floor 取余(结果符号跟随除数),fmod是 C 风格截断取余(结果符号跟随被除数)。对浮点输入,mod通过a - floor_div(a,b)*b复合实现,fmod通过a - trunc_div(a,b)*b复合实现;对整型则分别落到FLOORMOD/CMOD。where是唯一的三元运算,底层直接映射Ops.WHERE;masked_fill就是where的封装:mask.where(value, self)。pow的整数限制:当两个输入都是整数时,除int且非负指数外会抛RuntimeError("base needs to be float"),因为底层POW原语按浮点语义实现。- 文档中
add/sub/mul/div等方法的 docstring 都明确声明支持"broadcasting to a common shape, type promotion, and integer, float, boolean inputs",这些语义正是由_broadcasted保证的。
示例:
from tinygrad import Tensor t = Tensor([[1., 2.], [3., 4.]]) print((t + 10).numpy()) # 标量广播: [[11. 12.] [13. 14.]] print((t * Tensor([[2.], [0.5]])).numpy()) # 张量广播 print(Tensor([-4, 7, 5]).mod(Tensor([2, -3, 8])).numpy()) # floor 取余 print(Tensor([-4, 7, 5]).fmod(Tensor([2, -3, 8])).numpy()) # C 风格取余 print(Tensor([True, False]).where(1, 3).numpy()) # [1 3] print(Tensor([-1., 2.]).logaddexp(Tensor([-2., 3.])).numpy())Casting Ops(类型转换)
类型转换操作定义在 tinygrad/mixin/dtype.py 中,底层对应Ops.CAST与Ops.BITCAST:
| 方法 | 底层操作 | 说明 |
|---|---|---|
cast(dtype) | CAST | 数值类型转换(int↔float 会做值转换) |
bitcast(dtype) | BITCAST | 按位重解释(不改变底层比特,仅改变解释方式) |
float() | cast(float32) | 转单精度浮点 |
half() | cast(float16) | 转半精度浮点 |
int() | cast(int32) | 转 32 位整型 |
bool() | cast(bool) | 转布尔(非零为 True) |
bfloat16() | cast(bfloat16) | 转 bfloat16 |
double() | cast(double) | 转双精度浮点 |
long() | cast(long) | 转长整型 |
short() | cast(short) | 转短整型 |
cast与bitcast的本质区别:cast改变数值含义(如float32 → int32会截断取整),bitcast仅改变对同一组比特的解释(如把float32的比特重解释为int32得到的是该浮点数的 IEEE 754 位模式)。这也是 tinygrad/uop/init.py 注释中特别提到BITCAST是否属于 Elementwise 尚有讨论空间的原因——它可能伴随形状/布局变化。
另外值得注意的是 tinygrad/mixin/elementwise.py 的ufix机制:当一个 Python 标量参与运算时,它会被包装成弱类型常量(weak const),并在类型提升过程中保持"弱"属性,避免过早固定宽度导致精度损失——这是 tinygrad 编译器做常量折叠优化时的关键设计。
运算符重载速查
ElementwiseMixin为大多数二元运算提供了 Python 运算符重载(见 tinygrad/mixin/elementwise.py),因此可以直接写表达式:
from tinygrad import Tensor a = Tensor([1., 2., 3.]) b = Tensor([4., 5., 6.]) c = a + b # add d = a - b # sub e = a * b # mul f = a / b # div g = a // b # div(rounding_mode="floor") h = a % b # mod i = a ** 2 # pow j = 10 - a # __rsub__(reverse) k = ~Tensor([0, 1], dtype="int8") # bitwise_not m = a < b # __lt__ → CMPLT,结果是 bool 张量 n = a == b # eq(注意: 不重载 __eq__,等价于 Python 默认对象相等) o = a != b # ne值得注意的两个特例:
==没有重载:tinygrad/mixin/elementwise.py 注释明确说明__eq__保持 Python 默认语义(对象同一性),逐元素相等请用eq()方法;- 比较运算符:
<、>、<=、>=都基于底层CMPLT(小于)推导,<=即(>x).logical_not(),保证各后端只需实现一个比较原语。
实战:组合使用与验证
逐元素运算是构建更复杂结构的基本积木。例如实现 LayerNorm 前的归一化、自注意力中的softmax分母等,都可以用上述算子组合。仓库的测试集(如 test/backend/test_ops.py、test/backend/test_tensor.py)对每一类逐元素运算都做了 CPU 参考实现对照验证,覆盖了广播、类型提升、整数/浮点/布尔输入等组合场景,是阅读实现细节的最佳佐证。
一个综合示例(softmax 核心公式):
from tinygrad import Tensor def softmax(x: Tensor, axis: int = -1) -> Tensor: x = x - x.max(axis=axis, keepdim=True) # 数值稳定: 减去最大值 e = x.exp() return e / e.sum(axis=axis, keepdim=True) x = Tensor([[1., 2., 3.], [1., 2., 3.]]) print(softmax(x).numpy())其中x.exp()是逐元素运算,max/sum是归约运算(见 docs/tensor/ops.md 相关文档),/则依赖本文介绍的广播二元运算div。
小结
- tinygrad 的逐元素运算不改变张量形状,分为四类:数学一元运算(33 个)、激活函数一元运算(28 个)、广播二元/三元运算(18 个)、类型转换(11 个);
- 所有高层 API 最终都收敛到底层 Ops 枚举 的
Unary/Binary/Ternary原语,由_broadcasted(广播 + 类型提升)和_binop(统一入口)驱动,这使得各后端只需实现少量原语即可覆盖全部逐元素运算,并天然支持内核融合; - 大量"高级"函数(
sigmoid、tanh、gelu、acos等)是用基础原语组合而成,理解其复合方式有助于你写出能被编译器高效融合的自定义表达式; - 类型转换中
cast(数值转换)与bitcast(比特重解释)语义不同,务必区分使用。
如需继续深入,可进一步阅读 tinygrad/mixin/elementwise.py 的完整实现、tinygrad/uop/init.py 的算子枚举,以及 test/backend/test_ops.py 中逐元素运算的数值验证用例。
【免费下载链接】tinygradYou like pytorch? You like micrograd? You love tinygrad! ❤️项目地址: https://gitcode.com/GitHub_Trending/tiny/tinygrad
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考