Taichi RFC 解读:AOT 支持所有 SNode——SNode 树类型化与字段本地化的设计之路
【免费下载链接】taichiProductive, portable, and performant GPU programming in Python.项目地址: https://gitcode.com/GitHub_Trending/ta/taichi
导读
本文基于仓库中的设计文档 docs/rfcs/20220413-aot-for-all-snode.md(Taichi 官方 RFC,作者 Ye Kuang,2022-04-13),系统解读 Taichi 如何通过"SNode 树类型(SNodeTree type)"这一抽象,让 AOT(Ahead-of-Time)编译支持任意类型的 SNode 与 Taichi 字段,从而把"全局变量式"的 Taichi 字段改造为可显式传入 kernel 的局部化实体。读完本文,你将理解:为什么ti.field()的全局化实现会成为 AOT 部署的瓶颈;RFC 提出的SNodeTreeBuilder(仓库中落地为FieldsBuilder)如何实现"类型构建与实例化解耦";shape、AoS/SoA、梯度字段、Python/C++ AOT API 分别如何设计;以及这套设计与仓库现有实现(FieldsBuilder、AOT Module、C++ 侧 module_loader.h 等)之间的对应关系。
背景:为什么"全部 SNode 都能 AOT"是个问题
在 RFC 写作时的 Taichi 中,字段的典型定义与使用方式如下:
a = ti.field(ti.i32) b = ti.field(ti.f32) ti.root.pointer(ti.ij, 16).dense(ti.ij, 16).place(a, b) @ti.kernel def run(): for I in ti.grouped(a): b[I] = a[I] * 4.2这种写法对 Python 用户非常友好,但对"部署侧"(AOT 场景)提出了三个挑战:
Taichi 字段目前是全局变量实现的。这导致 Taichi kernel 变得"不纯"(not pure),依赖隐式信息。将这样的 kernel 保存进 AOT 模块时,还必须把其依赖的全部全局状态一并保存。理想情况下,用户应该能创建 Taichi 字段,并像参数一样把它们传入 kernel。
AOT 模块中缺少 SNode 类型信息。要朝"把字段传入 kernel"的方向前进,字段与 SNode 的类型都必须被保存进 AOT 模块。
字段数据不由用户管理。由于字段是全局的,Taichi 运行时必须负责创建和管理它们。若把字段局部化、与 Taichi kernel 解耦,用户就能自行管理这些字段的内存资源。
RFC 由此给出了明确的Goals:
- 提供一种 SNode API,让 SNode 与 Taichi 字段可以被"局部化",从而让 kernel 变得纯(pure);
- 支持显式描述完整的 SNode 树类型;
- 使 SNode 类型可被序列化进 AOT 模块,从而让 AOT 支持所有种类的 SNode;
- 新 SNode API 需兼容既有用法;
- (不确定但强烈期望)将元素类型与 SNode 类型解耦,解决矩阵字段必须以"分散"方式实现才能支持 SoA 布局的问题。
同时明确了一个Non-Goal:不打算把稀疏 SNode 的支持从 LLVM codegen 扩展到其他后端(尤其是 SPIR-V)。
事实核对:上述三点背景与目标来自 RFC 原文;仓库实现侧,LLVM 后端 AOT builder 的注释也印证了"序列化最小单元是整棵 SNodeTree"的结论,见下文"仓库中的落地佐证"。
核心设计(一):第一次尝试为何行不通
一个直觉上的方案是允许字段作为 kernel 参数:
a = ti.field(ti.i32) b = ti.field(ti.f32) ti.root.pointer(ti.ij, 16).dense(ti.ij, 16).place(a, b) @ti.kernel def run(a: ?, b: ?): for I in ti.grouped(a): b[I] = a[I] * 4.2 run(a, b)但 RFC 明确指出:这对 AOT 并不真正可行,因为a和b是"一个树类型的属性"(attributes of a tree type),你无法单独 dumpa和b的类型。
为了讲清这个问题,RFC 用 C++ 做了等价类比:
struct AB { int32_t a; float b; }; using TreeType = PointerDense<AB>;此时你无法把 kernel 声明成void run(? a, ? b);正确做法是把整个TreeType实例作为一个整体传入,即void run(TreeType &tree)。
这背后的原因是:在使用 Taichi 的 SNode 系统构造层级结构的同时,你也在构造一个SNodeTree类型——该工作由 Taichi 的 FieldsBuilder 完成(RFC 原文此处即引用了该实现文件)。
核心设计(二):可行方案——类型与实例解耦
RFC 的解决思路是:显式化 SNode 树及其类型,引入SNodeTreeBuilder。每个字段通过add_field()注册到 builder 中;add_field()不做任何内存分配,只返回一个field handle(字段句柄),供 kernel 内部从树中取回字段。
builder = ti.SNodeTreeBuilder() builder.add_field(dtype=ti.f32, name='x') builder.add_field(dtype=ti.i32, name='y') builder.tree() .pointer(ti.ij, 4) .dense(ti.ij, 5) .place('x', 'y') # `tree_t` stands for "tree type". tree_t = builder.build()同理,SNodeTreeBuilder.build()也不为树分配内存,它只构建一棵 SNode 树的类型。之后你可以用tree_t.instantiate()来实例化一棵树。类型-树解耦的设计动机有两点:
- 我们显式拿到了 SNode 树类型。这对 AOT 是必须的,同时也可用作类型注解,提升语言的形式化程度。
- 我们可以从同一个类型实例化出任意多棵树,并传给同一个 kernel 而无需重新编译。
在 Taichi kernel 内部,整棵树可以这样使用:
@ti.kernel def run(tr: tree_t): for I in ti.grouped(tr.x): tr.x[I] = tr.y[I] + 2.0 tree = tree_t.instantiate() run(tree)与既有 API 的唯一变化是:字段前需要加上tree.前缀;下标操作仍发生在字段上而非树上(即tr.x[I],而不是tr[I].x)。
两种从树中取回字段的方式
按名称(by name):
add_field()接收name参数。构建完 SNode 树后,Taichi 会为该树上的每个已注册字段生成一个属性,因此可以直接写tr.x访问名为'x'的字段。name是字段在树中的唯一标识符;注意在place时传入的也是名字。按字段句柄(by field handle):也可以使用
add_field()返回的句柄来访问字段:builder = ti.SNodeTreeBuilder() x_handle = builder.add_field(dtype=ti.f32, name='x') # boilerplate to generate tree type and instantiate a tree ... @ti.kernel def foo(tr: tree_t): x = ti.static(tr.get_field(x_handle)) # 1 for i in x: x[i] = i * 2.0注意该设计要求 kernel 中的部分(第 1 行)在 Python 侧求值,同时把全局变量
x_handle拉进了 kernel,某种程度上违背了最初"纯化"的目标。RFC 对此的取舍是:可以要求x_handle作为参数传入 kernel,或者干脆把它看作一个无足轻重的 Python 常量。
定义shape
与ti.field()类似,add_field可以接收shape参数。一旦指定,builder 会自动在树根下创建一个新的dense字段;注意指定shape后就不应再做一次place:
builder = ti.SNodeTreeBuilder() builder.add_field(dtype=ti.f32, name='x', shape=(4, 8)) # This would result an error # builder.tree().dense(ti.ij, (4, 8)).place('x') tree_t = builder.build()它等价于显式写法:
builder = ti.SNodeTreeBuilder() builder.add_field(dtype=ti.f32, name='x') builder.tree().dense(ti.ij, (4, 8)).place('x') tree_t = builder.build()AoS 与 SoA:复合类型与字段视图(field view)
需要在 AoS/SoA 之间切换的两种复合类型是ti.Matrix与ti.Struct。
AoS 很直接:直接把复合类型用作字段的dtype即可。
builder = ti.SNodeTreeBuilder() builder.add_field(dtype=ti.vec3, name='x') # ti.vec3 is a vector of 3 ti.f32's builder.dense(ti.i, 8).place('x') tree_t = builder.build()SoA 则麻烦一些。RFC 写作时的现行做法是把复合类型的每个分量当作独立的标量 Taichi 字段:如下例,必须手动分别 placex的 3 个底层分量:
# Current way (as of v1.0.1) of doing SoA in Taichi x = ti.Vector.field(3, ti.f32) for f in x._get_field_members(): # `x` consists three scalar f32 fields ti.root.dense(ti.ij).place(f)这种做法在多处引入混乱:
- 类型不单纯由
dtype决定,还取决于字段如何被 place; - 引入了"嵌套字段"(nested field)概念,而 Taichi 对此缺乏良好抽象。这使得对复合类型字段做某些优化(例如在特定平台上向量化 load/save 与标量操作带宽相同)变得复杂——没有良好抽象时,判断矩阵字段是 AoS 还是 SoA 的检查不得不散布在 CHI IR 的不同 pass 中;
- 进一步思考会发现,SoA 的
x其实不是一个真正的字段,而是三个独立标量字段的分组视图(grouped view)——该视图提供对单个标量字段无意义的矩阵运算。
由于类型目前与字段定义耦合,Taichi 字段为了支持 SoA 场景不得不实现为一个个独立字段;一旦切换到类型 builder 模式,就可以先控制类型如何构建,再选择字段实现方式。
若想把"这是一个字段视图"显式表达出来,RFC 给出了add_field_view设计:
builder = ti.SNodeTreeBuilder() builder.add_field(dtype=ti.f32, name='v0') builder.add_field(dtype=ti.f32, name='v1') builder.add_field(dtype=ti.f32, name='v2') for v in ['v0', 'v1', 'v2']: builder.tree().dense(ti.ij, 4).place(v) # Checks that # 1. `components` and `dtype` are compatible. # 2. If `dtype` is a vector/matrix, then all the fields in `components` are homogeneous in their SNode hierarchy. builder.add_field_view(dtype=ti.vec3, name='vel', components=['v0', 'v1', 'v2'])矩阵字段视图支持常见的矩阵操作,等价于把每个分量展开成局部矩阵变量:
# 1 vel_soa[i, j].inverse() # equivalent to ti.vec3([v0[i, j], v1[i, j], v2[i, j]]).inverse() # 2 vel_soa[i, j][1] += 2.0 # equivalent to v1[i, j] += 2.0 # 3 vel_soa[i, j] = vel_soa[i, j] @ some_vec3 # equivalent to vel_tmp = ti.vec3([v0[i, j], v1[i, j], v2[i, j]]) vel_tmp = vel_tmp @ some_vec3 v0[i, j] = vel_tmp[0] v1[i, j] = vel_tmp[1] v2[i, j] = vel_tmp[2]字段视图还可以嵌套,例如用三个已注册字段构造出结构体视图:
vertex_t = ti.types.struct({'pos': ti.vec3, 'normal': ti.vec3}) sphere_t = ti.types.struct({'center': vertex_t, 'radius': ti.f32}) builder = ti.SNodeTreeBuilder() builder.add_field(dtype=ti.vec3, name='pos') builder.add_field(dtype=ti.vec3, name='normal') builder.add_field(dtype=ti.f32, name='radius') builder.add_field_view(dtype=sphere_t, name='spheres', components=[['pos', 'normal'], 'radius']) ### ^^^^^^^^^^^^^^^^^ Note this is nested梯度与自动微分
为支持 autodiff,add_field()仍需要接收needs_grad: bool参数:
b = ti.SNodeTreeBuilder() b.add_field(dtype=ti.f32, name='x', needs_grad=True) # AOS b.tree()....place('x', b.grad_of('x')) # or SOA b.tree()....place('x') b.tree()....place(b.grad_of('x'))当needs_grad=True时,原始(primal)字段与伴随(adjoint)字段定义在同一棵树内;需要用b.grad_of(primal_name)来获取伴随字段的句柄。RFC 特意指出,备选方案是使用f'{primal_name}.grad'这种命名约定,但"感觉太临时/太 hack"(too ad-hoc)。
如果你不想手动 place 梯度字段,也可以在末尾调用builder.lazy_grad(),它会自动 place 所有梯度字段。这一设计在仓库中确有对应实现:全局 builder 的lazy_grad会触发root.lazy_grad()(见 fields_builder.py),调试模式下还会在 materialize 时自动分配伴随 checkbit(见 impl.py 中root._allocate_adjoint_checkbit()的调用)。
Python AOT API:保存 SNode 树类型
RFC 设想的 Python AOT API 如下:
builder = ti.SNodeTreeBuilder() # ... tree_t = builder.build() @ti.kernel def foo(tr: tree_t): # ... m = ti.aot.Module(arch) m.add_snode_tree_type(tree_t, name="vel_tree") m.add_kernel(foo) m.save('/path/to/module')在仓库当前的落地实现中(见 python/taichi/aot/module.py),ti.aot.Module(arch)构造时会通过rtm._finalize_root_fb_for_aot()把全局根 FieldsBuilder 以"仅编译类型"(compile_only)的方式 finalize,然后由prog.make_aot_module_builder(arch, caps)创建后端对应的 builder;字段通过Module.add_field(name, field)加入(内部调用self._aot_builder.add_field(...)),kernel 通过Module.add_kernel(kernel_fn)加入(内部调用self._aot_builder.add(kernel_name, kernel.kernel_cpp)),最后Module.save(filepath)落盘,并在目录中额外写入__content__与__version__文件记录模块内容清单与 Taichi 版本。可见 RFC 中"整棵树类型入库"的思想在落地时演化为"以字段(其背后是整棵 SNodeTree)为单位入库",但"kernel 与字段类型分离、可独立加载"的架构与 RFC 一脉相承。
C++ AOT API:加载并实例化树
RFC 设想的 C++ 侧 API 清晰地演示了"按类型取树 → 分配内存 → 实例化 → 启动 kernel"的完整链路:
auto mod = taichi::aot::Module("/path/to/module"); auto *tree_t = mod->get_snode_tree("vel_tree"); taichi::Device::AllocParams alloc_params; alloc_params.size = tree_t->get_size(); auto *tree_mem = device->allocate_memory(alloc_params); // By doing this, the kernel can verify that the passed in memory matches its // signature. auto *tree = taichi::instantiate_tree(tree_t, tree_mem); auto foo_kernel = mod->get_kernel("foo"); foo_kernel->launch(/*args=*/{tree});关键点在于:内存由用户(宿主程序)分配,kernel 在启动时可校验传入的内存与其签名是否匹配。这与背景中"字段数据不再由 Taichi 运行时管理"的目标直接呼应。
仓库中aot::Module确实提供了Field *get_snode_tree(const std::string &name)接口(见 taichi/aot/module_loader.h),Field类还定义了ArgUnion = std::variant<bool, int64_t, uint64_t, const Field *>作为 kernel 参数联合类型(module_loader.h),说明"以整棵树作为 kernel 参数"已成为 AOT 加载侧的正式形态。
向后兼容:ti.root即全局 builder,ti.field()返回 thunk
RFC 要求新 API 兼容既有用法。当时的现状是:ti.root已经实现为一个"字段累加器"——root 中累积的所有字段会在 kernel 调用时被物化为一棵新的 SNode 树。
先看既有写法:
x = ti.field(ti.f32) ti.root.pointer(ti.i, 4).dense(ti.i, 8).place(x) @ti.kernel def foo(): for i in x: x[i] = i * 2.0其使用新 API 的等价写法为:
b = ti.SNodeTreeBuilder() b.add_field(ti.f32, name='x') b.tree().pointer(ti.i, 4).dense(ti.i, 8).place('x') tree_t = b.build() tr = tree_t.instantiate() @ti.kernel def foo(): for i in tr.x: tr.x[i] = i * 2.0为实现向后兼容,需要两类辅助机制:
- 把
x@old映射到tr.x@new,且运行时需要知道x@old属于哪棵 SNode 树; ti.field()返回的x@old在ti.root当前 SNode 树被构建并实例化之前,只是一个字段占位符。
RFC 给出的可行方案是:ti.root就是一个全局的SNodeTreeBuilder;ti.field()返回一个FieldThunk(thunk 即"延迟求值"的占位对象):
class FieldThunk: def __init__(self, fid): self.field_id = fid self.tree = None def bind(self, tree): self.tree = tree def field(dtype, name='', shape=None, offset=None, needs_grad=False): name = name or random_name() handle = ti.root.add_field(dtype, name) ft = FieldThunk(handle) ti.root._field_thunks.append(ft) return ft在物化 SNodeTree 时:
tree_t = ti.root.build() tree = tree_t.instantiate() ti._runtime.global_snode_trees.append(tree) for ft in ti.root._field_thunks: ft.bind(tree) # Make `ti.root` a new SNodeTreeBuilder to allow for dynamic fields ti.root = SNodeTreeBuilder()JIT 编译 Taichi kernel 时,把x@old变换为x.tree.get_field(x.field_id)(其中x是FieldThunk)。
仓库实现对照:这一"全局 root builder + 延迟 finalize + 重建新 builder"的模式在仓库中真实存在。Runtime维护unfinalized_fields_builder注册表(impl.py),materialize_root_fb()在首次 kernel 调用或 AOT 时 finalize 全局 root,并随后重建一个新的全局FieldsBuilder以支持动态字段(impl.py);未 finalize 的非 root builder 会在 kernel 编译前被validate_fields_builder()拦截报错。这与 RFC 的"每次物化后把ti.root换成新 builder"的设想一致。
仓库中的落地佐证:从 RFC 到实现
RFC 是 2022-04 的设计提案,其核心思想在仓库中已有相当程度的落地,可沿以下路径继续深入阅读:
字段构建器:python/taichi/_snode/fields_builder.py 中的
FieldsBuilder是 RFC 中SNodeTreeBuilder的落地形态(对外暴露为ti.FieldsBuilder与全局ti.root)。它提供dense/pointer/dynamic/bitmasked/quant_array/place/lazy_grad/finalize等接口;finalize(compile_only=False)与_finalize_for_aot()(即compile_only=True)分别对应"运行时物化"与"AOT 仅编译类型"两种路径(fields_builder.py)。注意:pointer、dynamic、bitmasked等稀疏类型在构造时会检查当前后端是否支持 sparse extension,不支持则抛出TaichiRuntimeError——这正是 RFC Non-Goal(稀疏 SNode 暂不扩展到 SPIR-V 等后端)在实现层的体现(fields_builder.py)。AOT 模块:python/taichi/aot/module.py 的
Module类负责把 kernel/字段/图序列化到磁盘目录,并支持.tcm归档打包(archive())。LLVM 后端的序列化粒度:taichi/runtime/llvm/llvm_aot_module_builder.cpp 的
add_field_per_backend()注释明确写道:"字段指 SNodeTree 中的叶子(Place SNode);单独序列化叶子或其分支没有意义,我们必须序列化的最小单元是整棵 SNodeTree;且 SNodeTree 以snode_tree_id作为标识符,而非字段名(多个字段可能指向同一棵 SNodeTree)。"这从实现层面印证了 RFC"无法单独 dump 字段类型、必须整体保存树类型"的核心论断。GFX 后端的树内存管理:taichi/runtime/gfx/snode_tree_manager.cpp 的
SNodeTreeManager通过materialize_snode_tree()编译 SNode 结构并分配 root buffer,通过get_field_in_tree_offset()计算树内字段偏移、get_snode_tree_device_ptr()取得设备指针——对应 RFC C++ API 中"实例化树并管理其内存"的职责划分。C++ 端测试验证:tests/cpp/aot/llvm/field_aot_test.cpp 展示了完整的 C++ 加载流程:
mod->get_kernel(...)取出 kernel、mod->get_snode_tree("0")按snode_tree_id取树、LLVM::allocate_aot_snode_tree_type()分配树内存,随后通过LaunchContextBuilder设置参数并依次 launchinit_fields、check_init_x等 kernel,覆盖 CPU(LlvmAotTest.CpuField)与 CUDA(LlvmAotTest.CudaField,在TI_WITH_CUDA且 CUDA 可用时运行)两个后端,还包含对 pointer 字段 deactivate/activate 的验证——即"AOT 支持全部 SNode(含稀疏)"在 LLVM 后端的回归测试。
备选方案与 FAQ
RFC 在 Alternatives 一节坦言:"不确定是否有更好的设计能覆盖上述全部目标"。FAQ 一节当时标注为 TBD(待补充),本文不臆造其内容。
小结
这条 RFC 的价值在于指出了 Taichi 从"Python 内嵌的全局字段 DSL"走向"可部署的 AOT 运行时"之间最关键的抽象缺口:字段类型无法脱离 SNode 树类型而独立存在。其给出的答案——引入显式的树类型构建器、类型与实例解耦、以整棵树为 AOT 序列化与 kernel 参数的最小单元、用FieldThunk兼容旧 API——在仓库的FieldsBuilder、Module、SNodeTreeManager与 LLVM/GFX AOT builder 中均有迹可循。对希望深入理解 Taichi AOT 工作流(tests/cpp/aot 目录下有大量相关测试)或在其上做二次开发的读者而言,这份 RFC 与上述源码共同构成了一条完整的学习路径。
【免费下载链接】taichiProductive, portable, and performant GPU programming in Python.项目地址: https://gitcode.com/GitHub_Trending/ta/taichi
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考