从一段 Triton 源码看访问者模式:编译器是如何"读懂"你的 Python 代码的

想要系统学习triton,可以参考笔者的github:https://github.com/shizhengLi/triton-learning

如果你写过 Triton kernel,可能享受过这样的体验:用纯 Python 写几行代码,加上一个 @triton.jit 装饰器,就能生成在 GPU 上高效运行的代码。但你有没有想过,Triton 是怎么把你写的 Python 函数"翻译"成 GPU 代码的?

答案藏在编译器的一个经典环节里——遍历抽象语法树(AST)。而 Triton 在这里用到的,正是软件工程中久经考验的访问者模式(Visitor Pattern)。

这篇文章我们就从 code_generator.py 里的一小段代码出发,把 Triton 的工作原理和访问者模式讲清楚。

先认识一下主角:那段代码

class ContainsReturnChecker(ast.NodeVisitor):
    """检查 AST 中是否包含早期返回语句"""

    def __init__(self, gscope):
        self.gscope = gscope

    def _visit_stmts(self, body) -> bool:
        return any(self.visit(s) for s in body)

    def visit_Return(self, node: ast.Return):
        return True

    def visit_If(self, node: ast.If):
        return (self._visit_stmts(node.body) or
                self._visit_stmts(node.orelse))

    def visit_For(self, node: ast.For):
        return self._visit_stmts(node.body)

    def visit_While(self, node: ast.While):
        return self._visit_stmts(node.body)

这个类干的事情很单纯:检查一段代码里有没有提前 return。但它的写法——继承 ast.NodeVisitor、定义一堆 visit_XXX 方法——就是访问者模式的标准长相。

要看懂它,我们得先聊聊 Triton 在做什么。

Triton 是什么,以及它的编译流程

Triton 是一个用来写 GPU kernel 的领域特定语言(DSL)。它最大的特点是:你用 Python 的语法写 kernel,但这些代码并不会真的当成普通 Python 跑。

当你给一个函数加上 @triton.jit 之后:

@triton.jit
def add_kernel(x_ptr, y_ptr, output_ptr, n_elements, BLOCK_SIZE: tl.constexpr):
    pid = tl.program_id(axis=0)
    block_start = pid * BLOCK_SIZE
    offsets = block_start + tl.arange(0, BLOCK_SIZE)
    mask = offsets < n_elements
    x = tl.load(x_ptr + offsets, mask=mask)
    y = tl.load(y_ptr + offsets, mask=mask)
    output = x + y
    tl.store(output_ptr + offsets, output, mask=mask)

Triton 会经历这样几个阶段,把它一路降级(lowering)成 GPU 能执行的机器码 [1]:

  1. Python 源码 → AST(抽象语法树)
  2. AST → Triton IR(基于 MLIR 的中间表示)
  3. Triton IR → Triton GPU IR
  4. → LLVM IR → PTX(NVIDIA)或对应后端 → 最终的 cubin/二进制

我们关心的第 1 到第 2 步,就发生在 code_generator.py 里。Triton 用 Python 自带的 ast 模块把源码解析成语法树,然后遍历这棵树,把每个节点翻译成对应的 IR 指令。

而"遍历语法树并对不同节点做不同处理"——这恰恰是访问者模式最经典的应用场景。

什么是抽象语法树(AST)

在讲访问者模式之前,先快速理解 AST。

当 Python 解析 output = x + y 这行代码时,它不会把它当成一个字符串,而是构造出一棵树:

Assign(赋值)
├── target: Name(id='output')
└── value: BinOp(二元运算)
            ├── left:  Name(id='x')
            ├── op:    Add(加号)
            └── right: Name(id='y')

整个函数体就是一堆这样的节点嵌套组合而成的树。节点的类型五花八门:ReturnIfForWhileAssignBinOpCall……

编译器的任务就是走遍这棵树,对每种节点做出恰当的反应。问题来了:怎么优雅地"对不同类型的节点做不同的事"?

访问者模式登场

最直白的写法可能是这样:

def handle(node):
    if isinstance(node, ast.Return):
        return True
    elif isinstance(node, ast.If):
        ...
    elif isinstance(node, ast.For):
        ...
    elif isinstance(node, ast.While):
        ...
    # 还有几十种节点类型……

这种一长串 if-elif-isinstance 的写法有几个明显的痛点:

  • 节点类型一多,这个函数就变成一坨难以维护的巨型分支。
  • 每加一种节点处理逻辑,就得改这个函数,违反开闭原则。
  • 不同的处理任务(检查 return、生成 IR、做类型推导)都得各自写一遍这套分支。

访问者模式就是用来解决这个问题的。它的核心思想是:把"如何处理某类节点"的逻辑,拆成一个个独立的方法,再用一套统一的分发机制,根据节点类型自动调用对应的方法。

Python 标准库的 ast.NodeVisitor

Python 在 ast 模块里已经内置了访问者模式的基础设施——ast.NodeVisitor。它的核心是一个 visit 方法,大致逻辑是这样:

class NodeVisitor:
    def visit(self, node):
        # 根据节点类名拼出方法名,比如 Return -> "visit_Return"
        method = 'visit_' + node.__class__.__name__
        visitor = getattr(self, method, self.generic_visit)
        return visitor(node)

    def generic_visit(self, node):
        # 默认行为:递归访问所有子节点
        for child in ast.iter_child_nodes(node):
            self.visit(child)

关键就在 visit 这几行。它拿到一个节点,看它是什么类型(比如 ast.Return),就拼出方法名 visit_Return,然后调用你定义的同名方法。如果你没定义对应方法,就退回到 generic_visit

这个"根据类型自动找方法"的过程,在访问者模式里叫做双重分发(double dispatch)。在静态类型语言(比如 C++/Java)里通常要靠 accept/visit 互相调用来实现,而 Python 借助 getattr 的动态特性,几行就搞定了。

回头再读那段代码

现在我们带着访问者模式的视角,重新看 ContainsReturnChecker

class ContainsReturnChecker(ast.NodeVisitor):

它继承 ast.NodeVisitor,自动获得了 visit 分发能力。这个类的目标很明确:判断一段代码里有没有 return 语句。

    def visit_Return(self, node: ast.Return):
        return True

遇到 Return 节点,直接返回 True——找到了,就是有 return。这是检查的"终点"。

    def _visit_stmts(self, body) -> bool:
        return any(self.visit(s) for s in body)

这是个辅助方法:给它一个语句列表(比如一个代码块),它逐条 visit,只要任意一条里发现了 return 就返回 Trueany 还自带短路特性,找到第一个就停。

    def visit_If(self, node: ast.If):
        return (self._visit_stmts(node.body) or
                self._visit_stmts(node.orelse))

遇到 if 语句,它要同时检查 if 分支体(node.body)和 else 分支体(node.orelse)。任何一个分支里有 return 都算数。

    def visit_For(self, node: ast.For):
        return self._visit_stmts(node.body)

    def visit_While(self, node: ast.While):
        return self._visit_stmts(node.body)

循环也一样,钻进循环体里继续找。

注意这里的精妙之处:这些方法是相互递归的。visit_If 调用 _visit_stmts,后者又对每条子语句调用 visitvisit 再分发到对应的 visit_Forvisit_Whilevisit_If……一层层往下钻。这样无论 return 藏在多深的嵌套结构里,都能被翻出来。

而对于那些不可能"包含" return 的节点(比如一个单纯的赋值语句 x = y),这个类压根没定义对应的 visit_XXX,于是落到 generic_visit 默认行为,自然返回 None(被 any 当成假值),检查继续。开发者只需要关心"哪些节点结构需要特殊处理",其余的交给基类,代码因此非常干净。

为什么 Triton 需要这个检查

你可能会问:检查有没有 return 这件事,对 GPU 编译有什么意义?

这跟 GPU 的执行模型有关。在 GPU 上,成千上万个线程是以 SIMT(单指令多线程)方式批量执行的,控制流并不像 CPU 那样可以让某个线程"提前跳出函数走人"。当一个 Triton kernel 里出现了早期返回——尤其是藏在 if 或循环里的 return——编译器就不能简单地按顺序生成代码,而需要采取特殊的控制流处理策略(比如插入条件谓词、改写控制流图)。

所以 Triton 在正式生成代码之前,先用 ContainsReturnChecker 快速扫一遍:这个函数体里到底有没有早期返回?根据结果决定走哪条代码生成路径。这是一个典型的"先分析、再决策"的编译器套路,而访问者模式让这个分析过程的代码既简洁又易扩展。

不止一个访问者

ContainsReturnChecker 只是 Triton 里诸多访问者中很小的一个。真正承担"把 AST 翻译成 IR"重任的是 CodeGenerator 类,它同样继承自 ast.NodeVisitor,但定义了数量多得多的 visit_XXX 方法:visit_BinOp 处理二元运算、visit_Call 处理函数调用、visit_For 把 Python 的 for 循环映射成 IR 的循环结构,等等。

这正体现了访问者模式的另一个优势:同一棵 AST,可以被多个不同的访问者从不同角度处理。一个负责检查 return,一个负责生成 IR,互不干扰,各自只关心自己那摊事。想新增一种分析(比如统计某种操作出现的次数),再写一个新的访问者类就行,完全不用动现有代码。

小结

回顾一下这趟旅程:

  • Triton 让你用 Python 写 GPU kernel,背后通过 @triton.jit 把源码编译成 GPU 代码,第一步就是把 Python 解析成 AST 并遍历它 [1]。
  • 遍历 AST 时,"对不同类型节点做不同处理"是个核心需求,访问者模式正是为此而生,它用一套 visit_XXX 方法 + 自动分发,取代了又臭又长的 if-isinstance 分支。
  • Python 标准库的 ast.NodeVisitor 借助 getattr 优雅地实现了这套分发机制。
  • ContainsReturnChecker 是个麻雀虽小五脏俱全的例子:通过递归地访问 If/For/While 的子语句,它能查出嵌套在任意深度的早期返回,帮助编译器选择正确的代码生成策略。

下次你写 Triton kernel 时,不妨想想:你敲下的每一行 Python,都会被解析成一棵树,然后被一群"访问者"轮番走访、逐一翻译。一个 1990 年代《设计模式》书里就写明白的经典套路,至今仍在最前沿的 GPU 编译器里默默工作着。好的设计,从来不会过时。


参考:Triton 官方仓库与编译流程文档 [1]

注:本文对 Triton 源码行为的描述基于 ast.NodeVisitor 的标准语义和公开的编译流程资料;code_generator.py 的具体实现细节会随 Triton 版本演进而变化,如需精确到某一版本的行为,建议直接对照对应 tag 的源码。

后记

2026年6月18日于上海,在claude opus 4.8辅助下完成。

更多推荐