Mojo 自定义类型运算符支持完全指南:通过 dunder 方法与 trait 为 struct 解锁完整运算符语法
【免费下载链接】mojoThe Modular Platform (includes MAX & Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo
本指南以 Mojo 官方手册《Operators for custom types》为核心,系统讲解如何通过实现 dunder 方法(如__add__、__radd__、__iadd__)以及遵循Equatable、Comparable、Boolable、Writable等 trait,让自定义 struct 获得与内置类型一致的运算符语法体验。文中以一个完整的Complex(复数)类型为贯穿案例,覆盖一元/二元/反向/就地运算、混合类型运算、相等比较、布尔上下文与下标访问等全部运算符类别,并结合仓库中的可运行测试与 Bazel 构建配置给出可直接验证的实践路径。
运算符机制总览:每个运算符对应一组 dunder 方法
Mojo 中的每个运算符都映射到一组可以在 struct 中实现的 dunder 方法(dunder,即 double underscore 双下划线方法)。只要在自定义类型中实现这些方法,就可以直接使用+、-、*、/、==、[]等运算符语法,而无需显式调用方法。
这种设计的核心价值在于:运算符语法打开后,自定义类型与内置类型在表达能力上对齐。例如实现了__add__()的Vector可以直接书写v1 + v2,而不是v1.add(v2)。
二元运算符的三种方法形态:forward / reverse / in-place
每个二元运算符最多对应三种方法形态。以加法a + b为例:
- 正向方法(Forward):Mojo 首先尝试调用
a.__add__(b); - 反向方法(Reverse):如果正向方法不存在,或无法处理
b的类型,Mojo 回退调用b.__radd__(a); - 就地方法(In-place):对于
a += b这种复合赋值,Mojo 调用a.__iadd__(b)。
反向方法专为混合类型表达式设计——当左侧操作数不认识右侧操作数的类型时发挥作用:
a + 5 # 调用 a.__add__(5) 5 + a # Int 不认识你的类型,回退到 a.__radd__(5) a += 5 # 调用 a.__iadd__(5)这一机制与 Python 的反射运算符语义一致,但 Mojo 在编译期完成方法分派,没有运行时开销。
一元运算符
一元运算符返回原值(若不变)或代表结果的新值。例如-x使用一元取反运算符:
@fieldwise_init struct MyInt: var value: Int def __neg__(self) -> Bool: return Self(-self.value)注:严格来说
__neg__应返回Self(即MyInt),而非Bool;此处修正为:
@fieldwise_init struct MyInt: var value: Int def __neg__(self) -> Self: return Self(-self.value)若x是MyInt实例,则-x返回一个value字段被取反的新实例。
比较运算符与 trait:Equatable 与 Comparable
运算符不强制要求类型 conform 到 trait——即使不实现任何 trait,只要定义了对应的 dunder 方法,运算符依然可用。但遵循 trait 有额外收益:默认实现让你少写代码,并获得类型系统的静态约束。
Comparabletrait 提供<=、>、>=的默认实现,你只需实现__lt__()和__eq__();Equatabletrait 在所有字段均为Equatable时,为__eq__()和__ne__()提供默认实现。
对于没有自然顺序的类型(如复数),只实现Equatable,不要实现Comparable——否则会错误地暗示一个复数"小于"另一个复数。
源码佐证:Equatable 的反射式默认实现
从仓库源码 comparable.mojo 可以看到,Equatable的默认__eq__()使用**编译期反射(comptime reflection)**逐字段比较:
# 默认实现使用反射:比较所有字段 comptime r = reflect[Self] comptime names = r.field_names() comptime types = r.field_types() comptime for i in range(names.length): comptime T = types[i] comptime assert conforms_to(T, Equatable), ... if r.field_refi != r.field_refi: return False return True默认__ne__()直接返回not self == other。这意味着:只要所有字段都 conform 到Equatable,简单的 struct 可以零方法实现相等比较。该源码文档还提醒两点边界:
- 默认实现执行成员级(memberwise)比较,对含浮点字段的类型需注意 NaN 语义(
NaN != NaN); - 对相互递归类型(如 struct A 含
List[B]字段、B 又含List[A]字段),反射式遍历会在编译期产生无限单态化循环导致编译器挂起,此时应至少为其中一个类型提供显式__eq__()。
源码佐证:Comparable 的默认比较方法
comparable.mojo 中Comparable(Equatable)要求实现__lt__()与__eq__(),其余运算符均由默认实现推导:
def __gt__(self, rhs: Self) -> Bool: # return rhs < self def __le__(self, rhs: Self) -> Bool: # return not rhs < self def __ge__(self, rhs: Self) -> Bool: # return not self < rhs源码注释同时指出:这些默认实现(通过__lt__推导)对比较开销大的类型可能低效,建议此类类型覆写全部默认实现。
下标运算符:getitem与setitem
实现__getitem__()支持读取、__setitem__()支持写入。两者都接受可变参数以支持多维索引。
一维简单集合,使用单一索引即可解锁下标语法:
struct MySeq[T: Copyable]: def __getitem__(self, idx: Int) -> T: ... def __setitem__(mut self, idx: Int, value: T): ...多维集合,可以使用多个索引参数或可变参数:
struct Grid[T: Copyable]: # 固定二维 def __getitem__(self, x: Int, y: Int) -> T: ... # 任意维度 def __getitem__(self, *indices: Int) -> T: ...支持切片:以 Slice 为参数
自定义下标还可以支持切片,如obj[1:5]。此时__getitem__()的参数类型改为Slice而不是Int。
每个Slice有三个可选字段:start、end、step。通过调用indices()并传入类型的尺寸来归一化:
struct MySeq[T: Copyable]: var size: Int def __getitem__(self, span: Slice) -> Self: var start: Int var end: Int var step: Int start, end, step = span.indices(self.size) ...indices()返回三元组,表示根据你的范围调整后的跨度,把省略值或负索引解析为非负位置。
实战演练:构建一个完整的 Complex 复数类型
接下来逐步构建一个Complexstruct。这个示例覆盖了每一类运算符实现:一元运算符、同类型与混合类型的二元运算符、反向方法、就地赋值、相等比较、布尔转换与下标访问。
说明:标准库中已包含参数化的复数类型
ComplexSIMD,支持基础算术;本示例中的Complex是独立实现,不基于ComplexSIMD,目的是完整演示运算符实现的所有形态。
第一步:创建基础类型
复数包含实部re与虚部im两个Float64字段:
from std.math import sqrt @fieldwise_init struct Complex( Boolable, Equatable, TrivialRegisterPassable, Writable, ): var re: Float64 var im: Float64TrivialRegisterPassable:赋予值语义(value semantics),无需编写特殊的生命周期方法;Equatable:允许比较两个实例;Writable:为print()语句生成输出;Boolable:允许Complex值用于布尔上下文,如if条件。
便捷初始化器
添加便捷初始化器后,可以用仅含实部的参数创建实例:
def __init__(out self, re: Float64): self.re = re self.im = 0.0第二步:让类型可打印(Writable 与 repr)
实现Writable后可直接使用print()和String()。自定义实现提供括号,并分别输出实部与虚部:
# Struct method def write_to(self, mut writer: Some[Writer]): writer.write("(", self.re) if self.im < 0: writer.write(" - ", -self.im) else: writer.write(" + ", self.im) writer.write("i)")还可以实现write_repr_to()定义值的表示形式(representation)——即repr()返回的开发者面向输出,通常镜像构造该值的代码形式:
# Struct method def write_repr_to(self, mut writer: Some[Writer]): t"Complex(re = {self.re}, im = {self.im})".write_to(writer)var c = Complex(3.14, -2.72) print(c) # (3.14 - 2.72i) print(repr(c)) # Complex(re = 3.14, im = -2.72)第三步:添加一元运算符支持
+c原样返回值;-c对两个分量取反:
# methods def __pos__(self) -> Self: return self def __neg__(self) -> Self: return Self(-self.re, -self.im) ... var c = Complex(-1.2, 6.5) print(+c) # (-1.2 + 6.5i) print(-c) # (1.2 - 6.5i)第四步:支持二元算术运算
为两个Complex值之间实现加、减、乘、除,每种形式返回新的Complex实例:
def __add__(self, rhs: Self) -> Self: return Self(self.re + rhs.re, self.im + rhs.im) def __sub__(self, rhs: Self) -> Self: return Self(self.re - rhs.re, self.im - rhs.im) def __mul__(self, rhs: Self) -> Self: return Self( self.re * rhs.re - self.im * rhs.im, self.re * rhs.im + self.im * rhs.re, ) def __truediv__(self, rhs: Self) -> Self: var denom = rhs.squared_norm() return Self( (self.re * rhs.re + self.im * rhs.im) / denom, (self.im * rhs.re - self.re * rhs.im) / denom, ) def squared_norm(self) -> Float64: return self.re * self.re + self.im * self.im def norm(self) -> Float64: return sqrt(self.squared_norm())说明:复数除法使用共轭技巧——
(a+bi)/(c+di) = (a+bi)(c-di)/(c²+d²),分母即rhs.squared_norm();norm()通过std.math.sqrt计算模长。
var c1 = Complex(-1.2, 6.5) var c2 = Complex(3.14, -2.72) print(c1 + c2) # (1.94 + 3.78i) print(c1 * c2) # (13.91 + 23.67i)第五步:用反向方法支持混合类型算术
要支持2.5 + c这种Float64在左的表达式,需要同时重载正向方法(Complex + Float64)与反向方法(Float64 + Complex)。没有__radd__()时2.5 + c会失败,因为Float64不认识Complex:
# Forward: Complex + Float64 def __add__(self, rhs: Float64) -> Self: return Self(self.re + rhs, self.im) # Reversed: Float64 + Complex def __radd__(self, lhs: Float64) -> Self: return Self(self.re + lhs, self.im) def __sub__(self, rhs: Float64) -> Self: return Self(self.re - rhs, self.im) def __rsub__(self, lhs: Float64) -> Self: return Self(lhs - self.re, -self.im) def __mul__(self, rhs: Float64) -> Self: return Self(self.re * rhs, self.im * rhs) def __rmul__(self, lhs: Float64) -> Self: return Self(lhs * self.re, lhs * self.im) def __truediv__(self, rhs: Float64) -> Self: return Self(self.re / rhs, self.im / rhs) def __rtruediv__(self, lhs: Float64) -> Self: var denom = self.squared_norm() return Self( (lhs * self.re) / denom, (-lhs * self.im) / denom, )现在两种操作数顺序都可用:
var c = Complex(-1.2, 6.5) print(c + 2.5) # (1.3 + 6.5i) print(2.5 + c) # (1.3 + 6.5i) print(2.5 * c) # (-3.0 + 16.25i)注意混合运算的不对称性:c - 2.5的__rsub__实现为lhs - self.re(把实部顺序颠倒),而2.5 / c的__rtruediv__需要分母使用squared_norm()——反向方法必须自行处理操作数顺序与除法分母,不能简单复用正向实现。
允许就地赋值(in-place)
就地方法直接修改self,而不是返回新值。可以为Complex与Float64两类操作数重载:
def __iadd__(mut self, rhs: Self): self.re += rhs.re self.im += rhs.im def __iadd__(mut self, rhs: Float64): self.re += rhs def __isub__(mut self, rhs: Self): self.re -= rhs.re self.im -= rhs.im def __isub__(mut self, rhs: Float64): self.re -= rhs def __imul__(mut self, rhs: Self): var new_re = self.re * rhs.re - self.im * rhs.im var new_im = self.re * rhs.im + self.im * rhs.re self.re = new_re self.im = new_im def __imul__(mut self, rhs: Float64): self.re *= rhs self.im *= rhs def __itruediv__(mut self, rhs: Self): var denom = rhs.squared_norm() var new_re = (self.re * rhs.re + self.im * rhs.im) / denom var new_im = (self.im * rhs.re - self.re * rhs.im) / denom self.re = new_re self.im = new_im def __itruediv__(mut self, rhs: Float64): self.re /= rhs self.im /= rhs ... var c = Complex(-1.0, -1.0) c += Complex(0.5, -0.5) print(c) # (-0.5 - 1.5i) c += 2.75 print(c) # (2.25 - 1.5i) c *= 0.75 print(c) # (1.6875 - 1.125i) c /= 2.0 print(c) # (0.84375 - 0.5625i)就地方法签名使用mut self(可变引用),区别于返回新值的普通二元方法;注意__imul__与__itruediv__先计算临时变量再写回,避免用尚未更新的字段参与运算。
第六步:支持类型相等比较
复数没有自然顺序,因此Complex遵循Equatable(而非Comparable),获得==与!=,同时不暗示"一个复数小于另一个"。
你不需要自己实现__eq__()或__ne__()——只有当类型需要与成员级字段比较不同的相等语义时才需手写。Equatable提供基于编译期反射的默认__eq__()逐字段比较,以及返回其反值的默认__ne__()。Complex相等当且仅当两个字段都相等,因此反射式默认实现正是所需行为。(注意:由于浮点NaN永不自等,包含NaN的Complex也不会等于自身。)
var c1 = Complex(-1.2, 6.5) var c2 = Complex(-1.2, 6.5) var c3 = Complex(3.14, -2.72) print(c1 == c2) # True print(c1 != c3) # True第七步:支持布尔上下文(Boolable +bool)
遵循Boolable并实现__bool__(),即可在if条件或Bool()调用中直接使用Complex值。Mojo 将内置数值视为"非零即真",因此自然的定义是:任一分量非零即视为真:
def __bool__(self) -> Bool: return self.re != 0.0 or self.im != 0.0var c1 = Complex(0.0, 0.0) var c2 = Complex(-1.2, 6.5) print(Bool(c1)) # False print(Bool(c2)) # True if c2: print("c2 is nonzero") # c2 is nonzero从源码 bool.mojo 可确认Boolable要求实现__bool__(),用于if/while条件与显式Bool转换。
第八步:解锁下标访问(getitem/setitem)
get 与 set item dunder 允许索引类型内部内容。本示例中实部是索引 0,索引 1 返回虚部:
def __getitem__(self, idx: Int) raises -> Float64: if idx == 0: return self.re if idx == 1: return self.im raise "index out of bounds" def __setitem__(mut self, idx: Int, value: Float64) raises: if idx == 0: self.re = value elif idx == 1: self.im = value else: raise "index out of bounds" ... var c = Complex(3.14) print(c[0], c[1]) # 3.14 0.0 c[1] = 42.0 print(c) # (3.14 + 42.0i)越界时通过raise "index out of bounds"抛出错误,因此两个方法签名带raises效果。
一个演练,覆盖全部运算符
本示例从简单算术到比较再到下标,覆盖了 Mojo 的每一类运算符。实现正确的 dunder 方法和/或遵循正确的 trait,就能让几乎任何自定义类型使用运算符语法。
配套代码与测试:立即运行验证
本指南对应的可运行代码与测试位于仓库 Mojo/docs/site/code/manual/structs/operator-support 目录:
- tests.mojo 是完整的独立 Mojo 应用程序,将本文所有代码片段整合为可编译、可运行、可断言的完整程序;
- BUILD.bazel 定义了构建与测试规则。
构建与测试配置解读
该 BUILD.bazel 采用列表推导式,对目录下每个.mojo文件自动生成两类目标:
load("//bazel:api.bzl", "modular_run_binary_test", "mojo_binary") MOJO_SRCS = glob(["*.mojo"]) [ mojo_binary( name = src.split(".")[0], srcs = [src], deps = [ "@mojo//:std", ], ) for src in MOJO_SRCS ] [ modular_run_binary_test( name = src.split(".")[0] + "_test", size = "small", binary = src.split(".")[0], ) for src in MOJO_SRCS ]- 每个
.mojo文件生成一个mojo_binary目标,命名取文件名去扩展名(即tests); - 依赖
@mojo//:std(Mojo 标准库),测试文件因此可使用std.testing与std.math; - 每个 binary 对应一个
modular_run_binary_test测试目标(tests_test,size = "small"),运行时可执行程序并校验退出状态。
测试覆盖清单
tests.mojo 的test_complex()使用std.testing的assert_equal/assert_true与std.math的sqrt/isclose断言:
- 字段访问与打印:
c.re/c.im值、String(c)输出(3.14 - 2.72i); - 一元运算符:
+c、-c的字符串表示; - 模长与范数:
squared_norm()(43.69)与norm()(6.6098)经isclose验证; - 同类型二元运算:
c1 + c2、c1 - c2的实虚部结果; - 混合类型运算:
c + 2.5、2.5 + c、c * 2.5等正/反向方法; - 就地运算链:
+=、-=、*=、/=对Complex与Float64两种操作数的连续运算; - 相等比较:
c1 == c2为真、c1 != c3为真; - 下标访问:
c[0]、c[1]读取与c[1] = 42.0写入。
main()仅调用test_complex(),即该文件既是示例应用又是自测程序。运行测试可通过 Bazel:
bazel test //Mojo/docs/site/code/manual/structs/operator-support:tests_test或直接运行示例程序:
bazel run //Mojo/docs/site/code/manual/structs/operator-support:tests总结:运算符实现决策速查
| 想要支持的语法 | 需要实现的 dunder / trait |
|---|---|
-x、+x | __neg__()、__pos__() |
a + b/a - b/a * b/a / b | __add__、__sub__、__mul__、__truediv__ |
混合类型(标量在左,如2.5 + c) | 额外实现__radd__、__rsub__、__rmul__、__rtruediv__ |
a += b等复合赋值 | __iadd__、__isub__、__imul__、__itruediv__(签名mut self) |
a == b/a != b | 遵循Equatable(默认反射实现)或手写__eq__/__ne__ |
a < b、<=、>、>= | 遵循Comparable,实现__lt__与__eq__ |
x[i]读写 | __getitem__()、__setitem__()(支持多索引/可变参数) |
obj[1:5]切片 | __getitem__(span: Slice),配合span.indices(size)归一化 |
if x:/Bool(x) | 遵循Boolable,实现__bool__() |
print(x)/repr(x) | 遵循Writable,实现write_to()/write_repr_to() |
核心要点总结:
- 二元运算符三形态:正向
__add__、反向__radd__、就地__iadd__,覆盖所有组合与复合赋值场景; - trait 是省力杠杆:
Equatable的反射默认实现与Comparable的推导默认实现能显著减少样板代码;但浮点 NaN 语义、递归类型等边界场景需要手写覆写; - 混合类型务必成对实现:反向方法不是可选项,左侧是内置类型时它是唯一让表达式成立的手段;
- 没有自然顺序就别实现 Comparable:语义正确性优先于方法数量;
- 下标与切片:
__getitem__/__setitem__支持多维索引;Slice参数配合indices()即可获得切片能力; - 测试即文档:仓库中 tests.mojo 与 BUILD.bazel 提供了完整可运行的验证闭环,是学习运算符实现的最佳参照。
【免费下载链接】mojoThe Modular Platform (includes MAX & Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考