Seyed Masoud Hosseini · Overview · Study log · Weekly summaries · Ideas · Search · Transcript · RSS feed

Machine Learning Compilation · Lecture 8 of 8 · 44:05

Ep 9: Computational Graph Optimization

Ep 9: Computational Graph Optimization on YouTube

Study guide

What this lecture covers

This lecture asks how to programmatically transform a high-level computational graph, rather than only doing so by hand as in earlier episodes. It uses a Relax program (a small IR for representing computational graphs of tensor operators) to show how the underlying AST is structured, then builds working passes using the visitor/mutator pattern.

By the end, you can write a pass that detects an operator pattern (such as multiply-then-add, or dense-then-add), rewrites it, and produces a transformed IR module - and you understand how MLC separates the decision of what to fuse from the later action of generating fused code. The lecture builds on the tensor function and framework-integration passes from previous episodes and sets up automatic scheduling covered afterward.

Key ideas

  • Relax function structure: a Relax function has parameters and a body that is a sequence expression containing data flow blocks, which in turn contain variable bindings whose right-hand side is a call node.
  • Visitor/mutator pattern: relax.PyExprMutator recursively visits every node in the AST; overriding visit_call_ lets a pass inspect a call node and optionally return a replacement node.
  • Pattern matching by hand: passes detect patterns (like fused multiply-add) by checking a node's operator, then looking up its bound inputs via the block builder to see if they match an expected earlier operation.
  • Dead code elimination: after rewriting, remove_all_unused deletes bindings that are no longer referenced, such as the original unfused operation.
  • Operator fusion into sub-functions: rather than inventing a new primitive operator for every possible fusion combination, the lecture's fusion pass groups matched operations (e.g. dense + add) into a separate sub-function tagged with a primitive attribute.
  • Lowering to TensorIR: a second pass walks the graph and replaces high-level ops (dense, add, relu) with calls to lower-level tensor functions via call_te, mirroring the earlier PyTorch importer.
  • fuse_tir pass: a built-in TVM pass that takes the grouped sub-functions and generates a single tensor function that actually stitches the separate computations into one, separating "what to fuse" from "how to fuse."
  • Compositionality: every transformation is an IR-module-to-IR-module pass, so fusion, lowering, and tensor-function transformations can be composed in different orders to build a full MLC pipeline.

Walkthrough

Recap and a first fusion example (0:00)

The lecture recaps prior transformations among tensor functions and introduces the goal: writing programs, not just performing manual edits, to transform high-level computational graphs. It sets up a small Relax program computing y + x * y and inspects its AST - parameters, a sequence expression, a data flow block, and variable bindings - to show what a pass will need to walk.

Writing a fused multiply-add pass with the visitor pattern (9:08)

Using relax.PyExprMutator, the lecture writes a pass that overrides visit_call_: it checks whether the current node is an add call, looks up its first input's binding to see if that in turn is a multiply call, and if so constructs a new ewise_fma call to replace both. Running visit_expr over the function replaces the matched pattern and, combined with remove_all_unused, produces a cleaner graph with the dead multiply binding removed.

Fusion operation (19:15)

The lecture applies the same technique to a more realistic case: fusing a dense (matrix multiply) layer followed by a bias add, using a small feedforward network built with the block builder API. Instead of inventing a new primitive dense_add operator (which would not scale to the many possible operator combinations), the pass extracts matched dense+add pairs into a separate sub-function marked with a primitive attribute, leaving a call to that sub-function in place of the original two operations.

Pattern Matching (25:15)

The matching logic is shown in detail: a generic match_call helper checks whether a node calls a given operator, first testing for an add call and then for a dense call feeding into it. Once matched, a new block builder constructs a standalone function with its own parameters that performs the dense-then-add computation, attaches the primitive attribute, adds it to the IR module, and replaces the original call site with a call to that new global function.

Mapping to tensor functions (30:18)

A second mutator lowers remaining high-level operators (dense, add, relu) to calls into lower-level tensor functions using call_te, guided by an operator map similar to the one used in the earlier framework-integration pass. Running this over the fused graph turns the grouped dense+add sub-function into two separate TensorIR-function calls inside it.

Fusing TensorIR into one function (20:03)

To actually merge the two separate tensor-function calls inside a fused sub-function into one combined tensor function, the lecture invokes the built-in fuse_tir pass rather than writing it by hand, since it is considerably more complex. Running the full pipeline - build model, fuse operators, lower to tensor functions, fuse TensorIR, then build and run - reproduces the original model's predictions.

Limitations and generalizing fusion (39:24)

The lecture closes by noting that this pattern-based fusion pass is deliberately simple: it can create redundant computation in some cases (such as a matmul feeding two separate adds) and only recognizes the specific patterns it is told to look for. It contrasts this with more general fusers that analyze tensor functions structurally (element-wise, broadcast, reduction) rather than by fixed operator names, and describes composing a custom fusion pass with a generic one as a practical middle ground.

Before you watch

  • Be familiar with the tensor function and IR module concepts from earlier MLC episodes, including how primitive tensor functions are represented.
  • Review the framework-integration lecture that imports a PyTorch model into Relax, since this lecture reuses its importer-style traversal pattern for lowering operators.
  • Familiarity with Python AST-style visitor patterns is helpful but not required; the lecture explains the mutator API as it goes.

Check your understanding

  1. What is the role of visit_call_ in a PyExprMutator, and how does a pass use it to detect and replace a matched pattern?
  2. Why does the fusion pass create a separate sub-function for a matched dense+add pattern instead of introducing a new primitive operator?
  3. What does the fuse_tir pass do that the earlier fusion pass, by itself, does not?
  4. Give an example the lecture gives of a fusion pattern that a naive pattern-based fuser would handle poorly, and explain why.

Vocabulary

programmatically (adverb)
Done automatically by a program, rather than by hand.
The lecture transforms the graph programmatically instead of by hand.
AST (noun)
Abstract syntax tree: a structured representation of code as a tree of parts.
The pass walks the AST of the Relax function.
sequence expression (noun)
A code structure that lists steps to be run one after another.
A Relax function body is a sequence expression of bindings.
binding (noun)
A statement that connects a variable name to a computed value.
Each variable binding stores the result of one operation.
visitor pattern (noun)
A programming technique that walks through every part of a structure and processes each one.
The visitor pattern lets a pass inspect every node in the graph.
mutator (noun)
A tool that walks through code and can change parts of it.
A PyExprMutator can replace call nodes with new ones.
override (verb)
To replace a general method with your own specific version.
The pass overrides visit_call_ to detect a pattern.
dead code elimination (noun)
Removing code whose results are never used.
Dead code elimination deletes the leftover unfused operation.
sub-function (noun)
A smaller function created inside a larger program to hold a specific piece of logic.
The matched operations are grouped into a sub-function.
attribute (noun)
A piece of extra information attached to an object or function.
The sub-function is tagged with a primitive attribute.
lower (verb)
To convert a high-level operation into a more detailed, low-level form.
A pass lowers high-level ops into calls to tensor functions.
redundant (adjective)
Repeated unnecessarily, wasting effort or resources.
Naive fusion can sometimes create redundant computation.
compositionality (noun)
The property that separate small pieces can be combined into a larger working whole.
Compositionality lets passes be combined in different orders.
structurally (adverb)
Based on the shape or pattern of the code, not just its names.
More general fusers analyze functions structurally, not by operator name.
middle ground (noun)
A balanced position between two different approaches.
Combining a custom pass with a generic one is a practical middle ground.
AST node (noun)
A single element within a program's tree-shaped structural representation.
The pass visits every AST node in the function.
detect (verb)
To notice or identify the presence of something.
The pass is written to detect a specific operator pattern.
rewrite (verb)
To change code into a new, equivalent form.
The mutator can rewrite a matched call into a new one.
global function (noun)
A function accessible from anywhere in a program, not limited to one scope.
The fused operations are added as a new global function.
traverse (verb)
To move through every part of a structure in order.
The visitor traverses the whole computation graph.

Chapters

← Ep 8: GPU and Specialized Hardware, Part 2