ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

tinygrad 逐元素(Elementwise)运算全解:一元数学、激活函数、广播二元运算与类型转换

2026/9/10 5:11:11 拓冰建站 浏览量
tinygrad 逐元素(Elementwise)运算全解:一元数学、激活函数、广播二元运算与类型转换 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, reverseFalse) - 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 定义了与逐元素运算直接对应的底层算子UnaryOpsCAST、BITCAST、EXP2、LOG2、SIN、SQRT、RECIPROCAL、NEG、TRUNCBinaryOpsADD、MUL、SHL、SHR、CDIV、MAX、CMOD、CMPLT、CMPNE、CMPEQ、XOR、OR、AND、THREEFRY、SUB、FDIV、POW、FLOORDIV、FLOORMODTernaryOpsWHERE、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.LOG2log10以 10 为底的对数log2() * log10(2)exp自然指数cast(float32).mul(1/ln2).exp2()最终落到Ops.EXP2exp22 的幂直接Ops.EXP2sqrt平方根直接Ops.SQRTrsqrt平方根倒数sqrt().reciprocal()sin/cos/tan三角函数sin直接Ops.SINcos用sin(π/2 - x)复合tan为sin/cosasin/acos/atan反三角函数asin用多项式逼近系数来自 Abramowitz Stegun 4.4.46acos π/2 - asinatan复合asintrunc向零截断直接Ops.TRUNCceil/floor向上/向下取整基于trunc与where复合round四舍五入银行家舍入实现保证 half-to-even 语义与 NumPy 一致isinf/isnan/isfinite数值状态检测isnan即self ! selfisfinite为二者取反lerp线性插值self (end - self) * weightuint8 有定点加速路径square平方self * selfclamp/clip数值截断clip是clamp的别名支持单边Nonesign符号函数基于where与常量abs绝对值self * self.sign()reciprocal倒数直接Ops.RECIPROCAL几个值得注意的实现细节cos不使用专门的Ops.COS。从源码可见cos通过sin(π/2 - x)计算这是为了减少后端需要实现的硬件原语数量让 CUDA/Metal/OpenCL 等各后端只需实现SIN一个三角函数即可覆盖cos、tan。round是四舍六入五成双bankers 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中大部分由基础算子组合而成没有引入新的硬件原语因此可以在任意后端上运行方法公式/含义默认参数实现要点relumax(x, 0)—(self 0).where(self, 0)sigmoid1 / (1 e^(-x))—基于exp2复合避免引入EXP原语logsigmoidlog(sigmoid(x))—-(-x).softplus()hardsigmoid分段线性 sigmoid 近似alpha1/6, beta0.5两个relu之差elu指数线性单元alpha1.0分段wherecelu连续可微 ELUalpha1.0alpha * elu(x/alpha)selu缩放 ELUalpha1.67326, gamma1.0507gamma * elu(alpha*x)swish/silux * sigmoid(x)—silu就是swish的别名relu6min(max(x,0),6)—relu().minimum(6)hardswishx * relu6(x3) / 6—三个算子的组合tanh双曲正切—2 * sigmoid(2x) - 1sinh/cosh双曲正弦/余弦—基于exp组合atanh/asinh/acosh反双曲函数—基于log、sqrt、square组合hardtanh硬双曲正切min_val-1, max_val1就是clip(min_val, max_val)erf误差函数—Abramowitz Stegun 7.1.26 多项式逼近gelu高斯误差线性单元approximatetanh支持tanh与none两种模式quick_gelux * sigmoid(1.702x)—Sigmoid GELU 近似leaky_relu带泄漏的 ReLUneg_slope0.01(self 0).where(neg_slope*self, self)mishx * tanh(softplus(x))—组合实现softpluslog(1 e^x)beta1.0(1/beta) * logaddexp(beta*x, 0)softsignx / (1 |x|)—组合实现源码中有两个值得注意的坑被显式注释出来relu不能用self.maximum(0)实现tinygrad/mixin/elementwise.py 注释说明maximum(0)在x 0处会产生错误的梯度会同时从两条路径回传一半梯度因此relu刻意写成(self 0).where(self, 0)以保证在 0 点处的梯度正确性。gelu默认使用 tanh 近似approximatetanh等价于 PyTorch 的默认approximatetanh传入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_slope0.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说明addADD加法sub-ADD对取反后的 b减法mul*MUL乘法div/FDIV复合支持rounding_modetrunc/floormod%FLOORMOD复合Python 风格 floor 取余fmod—CMOD复合C 风格截断取余符号跟随被除数bitwise_xor^XOR按位异或bitwise_andAND按位与bitwise_or\|OR按位或bitwise_not~复合按位非无符号取dtype.max异或有符号异或-1lshift/rshift/SHL/SHR算术移位要求整型pow**POW幂运算支持reverse如2.0 ** tmaximum—MAX逐元素最大值minimum—MAX复合逐元素最小值有符号整型用 XOR 技巧实现where—WHERE三元选择x_i if cond_i else y_icopysign—复合取self的幅值、other的符号logaddexp—复合数值稳定的log(e^a e^b)实现细节补充div的整数语义默认对整数做真除法先提升为浮点再除传入rounding_modetrunc或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.WHEREmasked_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)转布尔非零为 Truebfloat16()cast(bfloat16)转 bfloat16double()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_modefloor) h a % b # mod i a ** 2 # pow j 10 - a # __rsub__reverse k ~Tensor([0, 1], dtypeint8) # 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(axisaxis, keepdimTrue) # 数值稳定: 减去最大值 e x.exp() return e / e.sum(axisaxis, keepdimTrue) 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),仅供参考