Mojo 自定义类型运算符支持完全指南:通过 dunder 方法与 trait 为 struct 解锁完整运算符语法
2026/9/10 5:13:45 网站建设 项目流程

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__)以及遵循EquatableComparableBoolableWritable等 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)

xMyInt实例,则-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__推导)对比较开销大的类型可能低效,建议此类类型覆写全部默认实现。

下标运算符:getitemsetitem

实现__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有三个可选字段:startendstep。通过调用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: Float64
  • TrivialRegisterPassable:赋予值语义(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,而不是返回新值。可以为ComplexFloat64两类操作数重载:

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永不自等,包含NaNComplex也不会等于自身。)

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.0
var 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.testingstd.math
  • 每个 binary 对应一个modular_run_binary_test测试目标(tests_testsize = "small"),运行时可执行程序并校验退出状态。

测试覆盖清单

tests.mojo 的test_complex()使用std.testingassert_equal/assert_truestd.mathsqrt/isclose断言:

  • 字段访问与打印c.re/c.im值、String(c)输出(3.14 - 2.72i)
  • 一元运算符+c-c的字符串表示;
  • 模长与范数squared_norm()(43.69)与norm()(6.6098)经isclose验证;
  • 同类型二元运算c1 + c2c1 - c2的实虚部结果;
  • 混合类型运算c + 2.52.5 + cc * 2.5等正/反向方法;
  • 就地运算链+=-=*=/=ComplexFloat64两种操作数的连续运算;
  • 相等比较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()

核心要点总结:

  1. 二元运算符三形态:正向__add__、反向__radd__、就地__iadd__,覆盖所有组合与复合赋值场景;
  2. trait 是省力杠杆Equatable的反射默认实现与Comparable的推导默认实现能显著减少样板代码;但浮点 NaN 语义、递归类型等边界场景需要手写覆写;
  3. 混合类型务必成对实现:反向方法不是可选项,左侧是内置类型时它是唯一让表达式成立的手段;
  4. 没有自然顺序就别实现 Comparable:语义正确性优先于方法数量;
  5. 下标与切片__getitem__/__setitem__支持多维索引;Slice参数配合indices()即可获得切片能力;
  6. 测试即文档:仓库中 tests.mojo 与 BUILD.bazel 提供了完整可运行的验证闭环,是学习运算符实现的最佳参照。

【免费下载链接】mojoThe Modular Platform (includes MAX & Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询