首页
/ Taichi RFC 解读:AOT 支持所有 SNode——SNode 树类型化与字段本地化的设计之路

Taichi RFC 解读:AOT 支持所有 SNode——SNode 树类型化与字段本地化的设计之路

2026-09-10 16:59:18作者:韦蓉瑛

导读

本文基于仓库中的设计文档 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 分别如何设计;以及这套设计与仓库现有实现(FieldsBuilderAOT 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 场景)提出了三个挑战:

  1. Taichi 字段目前是全局变量实现的。 这导致 Taichi kernel 变得"不纯"(not pure),依赖隐式信息。将这样的 kernel 保存进 AOT 模块时,还必须把其依赖的全部全局状态一并保存。理想情况下,用户应该能创建 Taichi 字段,并像参数一样把它们传入 kernel。

  2. AOT 模块中缺少 SNode 类型信息。 要朝"把字段传入 kernel"的方向前进,字段与 SNode 的类型都必须被保存进 AOT 模块。

  3. 字段数据不由用户管理。 由于字段是全局的,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 并不真正可行,因为 ab 是"一个树类型的属性"(attributes of a tree type),你无法单独 dump ab 的类型。

为了讲清这个问题,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() 来实例化一棵树。类型-树解耦的设计动机有两点:

  1. 我们显式拿到了 SNode 树类型。这对 AOT 是必须的,同时也可用作类型注解,提升语言的形式化程度。
  2. 我们可以从同一个类型实例化出任意多棵树,并传给同一个 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.Matrixti.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 字段:如下例,必须手动分别 place x 的 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)

这种做法在多处引入混乱:

  1. 类型不单纯由 dtype 决定,还取决于字段如何被 place;
  2. 引入了"嵌套字段"(nested field)概念,而 Taichi 对此缺乏良好抽象。这使得对复合类型字段做某些优化(例如在特定平台上向量化 load/save 与标量操作带宽相同)变得复杂——没有良好抽象时,判断矩阵字段是 AoS 还是 SoA 的检查不得不散布在 CHI IR 的不同 pass 中;
  3. 进一步思考会发现,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.pyroot._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@oldti.root 当前 SNode 树被构建并实例化之前,只是一个字段占位符。

RFC 给出的可行方案是:ti.root 就是一个全局的 SNodeTreeBuilderti.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)(其中 xFieldThunk)。

仓库实现对照:这一"全局 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)。注意:pointerdynamicbitmasked 等稀疏类型在构造时会检查当前后端是否支持 sparse extension,不支持则抛出 TaichiRuntimeError——这正是 RFC Non-Goal(稀疏 SNode 暂不扩展到 SPIR-V 等后端)在实现层的体现(fields_builder.py)。

  • AOT 模块python/taichi/aot/module.pyModule 类负责把 kernel/字段/图序列化到磁盘目录,并支持 .tcm 归档打包(archive())。

  • LLVM 后端的序列化粒度taichi/runtime/llvm/llvm_aot_module_builder.cppadd_field_per_backend() 注释明确写道:"字段指 SNodeTree 中的叶子(Place SNode);单独序列化叶子或其分支没有意义,我们必须序列化的最小单元是整棵 SNodeTree;且 SNodeTree 以 snode_tree_id 作为标识符,而非字段名(多个字段可能指向同一棵 SNodeTree)。"这从实现层面印证了 RFC"无法单独 dump 字段类型、必须整体保存树类型"的核心论断。

  • GFX 后端的树内存管理taichi/runtime/gfx/snode_tree_manager.cppSNodeTreeManager 通过 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 设置参数并依次 launch init_fieldscheck_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——在仓库的 FieldsBuilderModuleSNodeTreeManager 与 LLVM/GFX AOT builder 中均有迹可循。对希望深入理解 Taichi AOT 工作流(tests/cpp/aot 目录下有大量相关测试)或在其上做二次开发的读者而言,这份 RFC 与上述源码共同构成了一条完整的学习路径。

热门项目推荐
相关项目推荐

项目优选

收起
ops-transformerops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
C++
1.15 K
2.77 K
kernelkernel
deepin linux kernel
C
34
18
docsdocs
暂无描述
Markdown
900
5.83 K
ops-nnops-nn
本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。
C++
929
1.85 K
pytorchpytorch
作为 Ascend for PyTorch 社区的核心组件,TorchNPU 是昇腾专为 PyTorch 打造的深度学习适配插件,使 PyTorch 框架能够直接调用昇腾 NPU,为开发者提供昇腾 AI 处理器的超强算力。
Python
860
1.36 K
jiuwenswarmjiuwenswarm
JiuwenSwarm 是一款基于openJiuwen开发的智能AI Agent,它能够将大语言模型的强大能力,通过你日常使用的各类通讯应用,直接延伸至你的指尖。
Python
3.94 K
1.03 K
ops-mathops-math
本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。
C++
1.37 K
1.47 K
kernelkernel
openEuler内核是openEuler操作系统的核心,既是系统性能与稳定性的基石,也是连接处理器、设备与服务的桥梁。
C
534
603
AscendNPU-IRAscendNPU-IR
AscendNPU-IR是基于MLIR(Multi-Level Intermediate Representation)构建的,面向昇腾亲和算子编译时使用的中间表示,提供昇腾完备表达能力,通过编译优化提升昇腾AI处理器计算效率,支持通过生态框架使能昇腾AI处理器与深度调优
C++
548
398
leetcodeleetcode
🔥LeetCode solutions in any programming language | 多种编程语言实现 LeetCode、《剑指 Offer(第 2 版)》、《程序员面试金典(第 6 版)》题解
Markdown
77
23